Course Content
Model Deployment for AI Engineers
4 sections · 10 lessons
Serving with FastAPI and Flask
A recommendation model goes live behind a small Flask app. Load testing looked fine — 40 requests per second, 120 ms median. Two weeks in, the on-call engineer is paged at 03:00: p99 latency has gone to 14 seconds and the container is being killed by the health checker. Nothing about the model changed.
The cause turns out to be four lines of code. The app loaded the model inside the request handler instead of at startup, so every request deserialised 180 MB from disk. Under light traffic the page cache hid it. Under real traffic, memory pressure evicted the cache and every request went to disk at once.
Serving a model is not a hard algorithmic problem. It is a collection of small structural decisions — where the model is loaded, whether a handler blocks the event loop, what happens when the input is malformed — where getting one wrong costs you an order of magnitude. This lesson works through those decisions in both Flask and FastAPI.
The four things every model server does
Strip away the framework and any inference endpoint is the same pipeline:
HTTP request → parse and validate the payload (reject garbage before touching the model) → preprocess into tensors (resize, tokenise, normalise) → forward pass (the only part people think about) → postprocess and serialise (softmax, argmax, label lookup, JSON)HTTP responseTiming a real image classifier on 4 vCPUs breaks down roughly like this:
| Stage | Time | Nature of the work |
|---|---|---|
| JSON/multipart parse | 2 ms | CPU, holds the Python GIL |
| Base64 decode + JPEG decode + resize | 18 ms | Mostly C libraries, releases the GIL |
| Forward pass | 80 ms | C++ kernels, releases the GIL |
| Softmax + label lookup + JSON encode | 3 ms | CPU, holds the GIL |
Note that the forward pass is 78% of the time and none of it is Python. That single fact determines almost every concurrency decision later in this lesson.
Flask: the smallest thing that works
Flask is a synchronous WSGI framework. One request occupies one worker until it finishes. That simplicity makes it a good place to see the structural rules clearly.
1from flask import Flask, request, jsonify2import numpy as np, joblib, time, logging34app = Flask(__name__)56# Loaded ONCE, at import time, before the first request arrives.7MODEL = joblib.load("churn_model.joblib")8FEATURES = ["tenure_months", "monthly_charges", "support_tickets", "contract_type"]910@app.get("/health")11def health():12 return jsonify(status="ok", model_version=MODEL.version), 2001314@app.post("/predict")15def predict():16 payload = request.get_json(silent=True)17 if payload is None:18 return jsonify(error="body must be JSON"), 4001920 missing = [f for f in FEATURES if f not in payload]21 if missing:22 return jsonify(error="missing features", fields=missing), 4222324 try:25 x = np.array([[float(payload[f]) for f in FEATURES]])26 except (TypeError, ValueError):27 return jsonify(error="all features must be numeric"), 4222829 t0 = time.perf_counter()30 proba = float(MODEL.predict_proba(x)[0, 1])31 latency_ms = (time.perf_counter() - t0) * 10003233 app.logger.info("prediction p=%.4f latency_ms=%.1f", proba, latency_ms)34 return jsonify(churn_probability=round(proba, 4),35 will_churn=proba > 0.5,36 latency_ms=round(latency_ms, 1)), 200Four things in that snippet are doing load-bearing work.
Module-level model load. The 03:00 page above is exactly what happens when this line moves inside the handler. Loading at import time means the cost is paid once per process, at startup, where it is visible.
Distinct status codes. 400 means "this is not JSON", 422 means "this is JSON but the contents are wrong", 500 means "we broke". Collapsing all of them into 500 makes your error-rate dashboard useless: you cannot tell a client sending bad data from your own model crashing, and you will page yourself for the former.
Type coercion is validation. float(payload[f]) is wrapped in a try block because {"tenure_months": "twelve"} is perfectly valid JSON. Without the guard, the exception surfaces as a 500 and a stack trace in your logs.
Latency measured around the model call only. If you time the whole handler you cannot tell slow inference from slow parsing.
Anything that costs more than a few milliseconds and does not depend on the request body belongs at startup, not in the handler.
Running Flask properly
The development server is single-threaded and explicitly not for production. Use Gunicorn:
gunicorn --workers 4 --threads 2 --timeout 60 --bind 0.0.0.0:8000 app:appEach worker is a separate process with its own copy of the model. Four workers with a 180 MB model is 720 MB of RSS before you serve a request — set your container memory limit with that arithmetic in mind, or the OOM killer will do it for you.
FastAPI: validation and concurrency as first-class features
FastAPI is ASGI-based and uses Pydantic for validation. The practical difference is that the twenty lines of manual checking above collapse into a type declaration, and you get the concurrency model needed for I/O-heavy endpoints.
1from fastapi import FastAPI, HTTPException, UploadFile, File2from pydantic import BaseModel, Field, field_validator3from contextlib import asynccontextmanager4from typing import Literal, List5import torch, io, time6from PIL import Image7import torchvision.transforms as T89STATE = {}1011@asynccontextmanager12async def lifespan(app: FastAPI):13 STATE["model"] = torch.jit.load("defectnet_traced.pt", map_location="cpu")14 STATE["model"].eval()15 torch.set_num_threads(1) # explained in the concurrency section16 STATE["labels"] = ["ok", "scratch", "dent", "discolour"]17 STATE["tfm"] = T.Compose([18 T.Resize((224, 224)),19 T.ToTensor(),20 T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),21 ])22 yield # application serves requests here23 STATE.clear() # shutdown2425app = FastAPI(title="Defect Classifier", version="2.1.0", lifespan=lifespan)The lifespan context manager is the modern replacement for startup and shutdown event handlers. Everything before yield runs once before the server accepts traffic; everything after runs on shutdown. Loading the model here rather than at module import has a real benefit: if the load fails, the process exits during startup and your orchestrator never routes traffic to it.
Pydantic models that reject bad data precisely
1class Reading(BaseModel):2 tenure_months: int = Field(..., ge=0, le=600)3 monthly_charges: float = Field(..., gt=0, le=10000)4 support_tickets: int = Field(..., ge=0, le=500)5 contract_type: Literal["monthly", "annual", "two_year"]67 @field_validator("monthly_charges")8 @classmethod9 def not_absurd(cls, v: float) -> float:10 if v > 2000:11 raise ValueError("monthly_charges above 2000 is outside training range")12 return v1314class BatchRequest(BaseModel):15 items: List[Reading] = Field(..., min_length=1, max_length=64)1617class Prediction(BaseModel):18 label: str19 confidence: float20 model_version: str2122@app.post("/predict", response_model=Prediction)23def predict(reading: Reading):24 ...Three things this buys you that hand-written checks rarely do.
The max_length=64 on the batch is a denial-of-service control, not a convenience. Without it, a client can post 500,000 items and your worker allocates until the container dies. Every list that reaches a model needs an upper bound.
The Literal type means contract_type can only be one of three strings. This matters because scikit-learn pipelines with one-hot encoders often silently encode an unseen category as all-zeros, producing a confident prediction from a feature vector the model never saw in training.
The ge/le bounds encode the training distribution. A model trained on tenures of 0–120 months will happily return a number for tenure 9,000 — extrapolating far outside the training range, with no signal that it is doing so. Rejecting at the boundary is honest; predicting is not.
Validation is not about politeness to clients; it is about refusing to make predictions the model has no basis for.
By default FastAPI returns 422 with a field-by-field explanation of what failed. That is genuinely useful to whoever is integrating against you, and it costs nothing.
An image endpoint
1MAX_BYTES = 5 * 1024 * 102423@app.post("/predict/image", response_model=Prediction)4def predict_image(file: UploadFile = File(...)):5 if file.content_type not in {"image/jpeg", "image/png"}:6 raise HTTPException(415, f"unsupported content type: {file.content_type}")78 raw = file.file.read(MAX_BYTES + 1)9 if len(raw) > MAX_BYTES:10 raise HTTPException(413, "image exceeds 5 MB")1112 try:13 img = Image.open(io.BytesIO(raw)).convert("RGB")14 except Exception:15 raise HTTPException(422, "file is not a decodable image")1617 x = STATE["tfm"](img).unsqueeze(0)18 t0 = time.perf_counter()19 with torch.no_grad():20 logits = STATE["model"](x)21 probs = torch.softmax(logits, dim=1)[0]22 idx = int(probs.argmax())2324 return Prediction(label=STATE["labels"][idx],25 confidence=round(float(probs[idx]), 4),26 model_version="2.1.0")Reading MAX_BYTES + 1 rather than the whole file is deliberate — it caps memory before you have committed to the allocation. Checking the size after reading everything means a 2 GB upload has already landed in RAM by the time you reject it.
torch.no_grad() is not optional. Without it PyTorch builds an autograd graph for every forward pass, which on a ResNet-50 costs roughly 30% extra latency and, worse, holds references to intermediate activations. Under sustained load that shows up as memory that climbs and never comes back down.
The concurrency decision that actually matters
In FastAPI, whether you write def or async def changes where your handler runs, and getting it backwards is the most common serious mistake in this stack.
| Declaration | Runs on | Blocking work inside is |
|---|---|---|
def predict(...) | A thread from the threadpool (default 40 threads) | Fine — it blocks only that thread |
async def predict(...) | The single event loop | Catastrophic — it blocks every other request |
Work the arithmetic. Suppose a forward pass takes 80 ms and 10 requests arrive simultaneously.
With async def and a synchronous model(x) call inside, the event loop is frozen for 80 ms per request. Requests are served strictly one after another: the first finishes at 80 ms, the tenth at 800 ms. Throughput is 12.5 requests per second and it cannot exceed that no matter how many cores the machine has. Health-check endpoints also queue behind the work — which is how a busy server gets marked unhealthy and restarted, making the problem worse.
With plain def, all 10 go to separate threads. PyTorch releases the GIL inside its C++ kernels, so on 4 cores roughly 4 run genuinely in parallel; the batch completes in about 240 ms rather than 800 ms, and the event loop stays free to accept new connections and answer health checks.
The rule is short: use async def only if the body contains await. If you are calling a synchronous model, use def. If you genuinely need async — say the handler also calls a feature store over HTTP — push the blocking part off the loop:
1import anyio, httpx23@app.post("/predict/enriched")4async def predict_enriched(req: Reading):5 async with httpx.AsyncClient(timeout=2.0) as client:6 r = await client.get(f"http://feature-store/user/{req.user_id}") # real await7 features = build(req, r.json())8 probs = await anyio.to_thread.run_sync(run_model, features) # off the loop9 return {"probability": probs}Threads inside the model, too
There is a second layer of parallelism that quietly fights the first. PyTorch defaults to using every available core for a single forward pass. On a 4-vCPU container:
| Configuration | Single-request latency | Throughput at saturation |
|---|---|---|
1 worker, torch.set_num_threads(4) | 80 ms | ~12.5 req/s |
4 workers, torch.set_num_threads(1) | 200 ms | ~20 req/s |
| 4 workers, threads left at default (4 each) | 90–400 ms, erratic | ~11 req/s |
The third row is the trap and it is the default configuration. Sixteen compute threads fighting over four cores means constant context switching and cache thrashing, so you get worse throughput and worse latency than either deliberate choice. Always set torch.set_num_threads(1) (or OMP_NUM_THREADS=1) when running multiple workers.
The choice between rows one and two is a real product decision: minimise latency for a user waiting on a screen, or maximise throughput for a batch job. There is no default answer, only a default mistake.
Errors, timeouts, and the middleware layer
Uncaught exceptions leak stack traces to clients — including file paths and sometimes model internals. Catch them centrally:
1from fastapi.responses import JSONResponse2from starlette.middleware.base import BaseHTTPMiddleware3import uuid, logging45log = logging.getLogger("serving")67@app.exception_handler(Exception)8async def unhandled(request, exc):9 rid = getattr(request.state, "request_id", "unknown")10 log.exception("unhandled error request_id=%s path=%s", rid, request.url.path)11 return JSONResponse(status_code=500,12 content={"error": "internal error", "request_id": rid})1314class Context(BaseHTTPMiddleware):15 async def dispatch(self, request, call_next):16 rid = request.headers.get("x-request-id") or str(uuid.uuid4())17 request.state.request_id = rid18 t0 = time.perf_counter()19 response = await call_next(request)20 ms = (time.perf_counter() - t0) * 100021 response.headers["x-request-id"] = rid22 response.headers["x-process-time-ms"] = f"{ms:.1f}"23 log.info("%s %s %d %.1fms rid=%s", request.method, request.url.path,24 response.status_code, ms, rid)25 return response2627app.add_middleware(Context)The client gets an opaque error plus a request ID; you get the full trace in your logs, findable by that ID. When a user reports "it failed around noon", the ID is the difference between a two-minute lookup and an afternoon of grepping.
Rate limiting
Inference is expensive per request, so a single misbehaving client can consume an entire replica. A token-bucket limiter keyed on API key handles both steady rate and short bursts:
1import time2from collections import defaultdict3from fastapi import Header45RATE, BURST = 10.0, 20 # 10 requests/second sustained, bursts up to 206_buckets = defaultdict(lambda: [BURST, time.monotonic()])78def check_rate(api_key: str = Header(..., alias="x-api-key")):9 tokens, last = _buckets[api_key]10 now = time.monotonic()11 tokens = min(BURST, tokens + (now - last) * RATE)12 if tokens < 1:13 retry = (1 - tokens) / RATE14 raise HTTPException(429, "rate limit exceeded",15 headers={"Retry-After": str(int(retry) + 1)})16 _buckets[api_key] = [tokens - 1, now]17 return api_key1819@app.post("/predict", dependencies=[Depends(check_rate)])20def predict(reading: Reading): ...Trace the arithmetic. A client idle for two seconds has accumulated min(20, tokens + 2 × 10) tokens, so it is back at the 20-token cap and can fire 20 requests instantly. After that it is throttled to one per 100 ms. That shape — burst then steady — matches how real clients behave far better than a hard "10 per second" window.
The important limitation: this dictionary lives in one process. With four Gunicorn workers your effective limit is 40 per second, not 10, and the counts reset on every deploy. For anything you actually need to enforce, move the bucket to Redis so all replicas share it.
Timeouts and streaming
An endpoint with no timeout will eventually hold a thread forever. Bound it:
1@app.post("/predict")2async def predict(reading: Reading):3 try:4 with anyio.fail_after(5.0):5 return await anyio.to_thread.run_sync(6 run_model, reading, abandon_on_cancel=True)7 except TimeoutError:8 raise HTTPException(504, "inference exceeded 5s budget")abandon_on_cancel=True is what makes the timeout fire. Without it, AnyIO will not cancel a call running in a worker thread: the request waits for the model to finish however long it takes, and the 504 never happens. With it, the client gets its answer at five seconds — but the abandoned thread keeps computing in the background, so this bounds the caller's wait, not your CPU usage.
For generative models where the full response takes many seconds, stream tokens as they are produced. Time-to-first-token drops from the full generation time to a few hundred milliseconds, which changes the perceived experience completely even though total time is unchanged:
1from fastapi.responses import StreamingResponse2import json34@app.post("/generate")5def generate(req: GenRequest):6 def token_stream():7 for tok in model.stream(req.prompt, max_tokens=req.max_tokens):8 yield f"data: {json.dumps({'token': tok})}\n\n"9 yield "data: [DONE]\n\n"10 return StreamingResponse(token_stream(), media_type="text/event-stream")Background tasks and model reloading
Work that the client does not need to wait for — writing the prediction to a warehouse, pushing a metric — belongs after the response:
1from fastapi import BackgroundTasks23@app.post("/predict")4def predict(reading: Reading, bg: BackgroundTasks):5 result = run_model(reading)6 bg.add_task(record_prediction, reading.model_dump(), result) # runs after response7 return resultThis is for cheap, fire-and-forget work only. A background task that fails does so silently, and a queue of them lives in process memory, so a restart loses everything pending. Anything that must not be lost belongs in a real queue.
Hot-reloading a model without dropping requests is a swap of a single reference:
1import threading2_lock = threading.Lock()34@app.post("/admin/reload")5def reload_model(version: str, _=Depends(require_admin)):6 new = torch.jit.load(f"/models/{version}/model.pt", map_location="cpu")7 new.eval()8 with torch.no_grad(): # prove it works first9 new(torch.zeros(1, 3, 224, 224))10 with _lock:11 STATE["model"], STATE["version"] = new, version # atomic swap12 return {"loaded": version}Load and smoke-test the new model before touching the live reference. If the file is corrupt you return an error and the running model is untouched. Doing the swap first and validating afterwards means a bad file takes the service down. Note the memory implication: both models are resident during the load, so a 180 MB model needs 360 MB of headroom.
Flask or FastAPI, and what it costs to choose wrong
| Flask | FastAPI | |
|---|---|---|
| Model | WSGI, synchronous | ASGI, async-capable |
| Validation | Hand-written or Marshmallow | Pydantic, from type hints |
| API docs | Add Flask-RESTX or write them | OpenAPI + Swagger UI automatically |
| Concurrent I/O calls | One thread each | Free with await |
| Streaming responses | Awkward under WSGI | Native |
| Best fit | Simple CPU-bound endpoint inside an existing Flask estate | New services, especially with I/O or streaming |
For a single synchronous predict endpoint the throughput difference is small — both are bound by the forward pass, not the framework. FastAPI wins on validation and generated documentation, and it wins decisively as soon as a request involves a database, a feature store, or streaming output.
What to check before you call a server done
The failures in this lesson share a shape: each one is invisible at low traffic and severe at high traffic. Load testing at your expected peak is the only thing that surfaces them, and it belongs before launch, not after the first incident.
Concretely, before shipping: confirm the model is loaded exactly once per process by logging at load time and counting the lines. Send a deliberately malformed payload and confirm you get a 422 with a useful message, not a 500 with a stack trace. Send a 100 MB file and confirm the rejection happens before the allocation. Run your peak expected concurrency for ten minutes and watch RSS — flat is correct, a steady climb means you are retaining something, and no_grad is the first suspect. Then check that a health-check request served during that load returns in single-digit milliseconds; if it does not, something in your handlers is blocking the event loop.
Those five checks take under an hour and catch the great majority of what pages people at 03:00.