Model Deployment for AI Engineers

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.

Where a blocking model call actually runsDeclared with plain def• FastAPI runs it in a worker thread pool• The event loop staysfree for new requests• Concurrency is capped by the pool size• The right shape for ablocking forward passDeclared with async def• The body runs on the event loop itself• One inference stalls every other request• Ten 80 ms requests finish at 800 ms• Right only when the body is awaited I/O
One keyword decides it: an async handler doing 80 ms of blocking model work serialises the whole process.

The four things every model server does

Strip away the framework and any inference endpoint is the same pipeline:

Text
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 response

Timing a real image classifier on 4 vCPUs breaks down roughly like this:

StageTimeNature of the work
JSON/multipart parse2 msCPU, holds the Python GIL
Base64 decode + JPEG decode + resize18 msMostly C libraries, releases the GIL
Forward pass80 msC++ kernels, releases the GIL
Softmax + label lookup + JSON encode3 msCPU, 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.

Python
from flask import Flask, request, jsonifyimport numpy as np, joblib, time, loggingapp = Flask(__name__)# Loaded ONCE, at import time, before the first request arrives.MODEL = joblib.load("churn_model.joblib")FEATURES = ["tenure_months", "monthly_charges", "support_tickets", "contract_type"]@app.get("/health")def health():    return jsonify(status="ok", model_version=MODEL.version), 200@app.post("/predict")def predict():    payload = request.get_json(silent=True)    if payload is None:        return jsonify(error="body must be JSON"), 400    missing = [f for f in FEATURES if f not in payload]    if missing:        return jsonify(error="missing features", fields=missing), 422    try:        x = np.array([[float(payload[f]) for f in FEATURES]])    except (TypeError, ValueError):        return jsonify(error="all features must be numeric"), 422    t0 = time.perf_counter()    proba = float(MODEL.predict_proba(x)[0, 1])    latency_ms = (time.perf_counter() - t0) * 1000    app.logger.info("prediction p=%.4f latency_ms=%.1f", proba, latency_ms)    return jsonify(churn_probability=round(proba, 4),                   will_churn=proba > 0.5,                   latency_ms=round(latency_ms, 1)), 200

Four 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:

Bash
gunicorn --workers 4 --threads 2 --timeout 60 --bind 0.0.0.0:8000 app:app

Each 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.

Python
from fastapi import FastAPI, HTTPException, UploadFile, Filefrom pydantic import BaseModel, Field, field_validatorfrom contextlib import asynccontextmanagerfrom typing import Literal, Listimport torch, io, timefrom PIL import Imageimport torchvision.transforms as TSTATE = {}@asynccontextmanagerasync def lifespan(app: FastAPI):    STATE["model"] = torch.jit.load("defectnet_traced.pt", map_location="cpu")    STATE["model"].eval()    torch.set_num_threads(1)          # explained in the concurrency section    STATE["labels"] = ["ok", "scratch", "dent", "discolour"]    STATE["tfm"] = T.Compose([        T.Resize((224, 224)),        T.ToTensor(),        T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),    ])    yield                              # application serves requests here    STATE.clear()                      # shutdownapp = 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

Python
class Reading(BaseModel):    tenure_months: int = Field(..., ge=0, le=600)    monthly_charges: float = Field(..., gt=0, le=10000)    support_tickets: int = Field(..., ge=0, le=500)    contract_type: Literal["monthly", "annual", "two_year"]    @field_validator("monthly_charges")    @classmethod    def not_absurd(cls, v: float) -> float:        if v > 2000:            raise ValueError("monthly_charges above 2000 is outside training range")        return vclass BatchRequest(BaseModel):    items: List[Reading] = Field(..., min_length=1, max_length=64)class Prediction(BaseModel):    label: str    confidence: float    model_version: str@app.post("/predict", response_model=Prediction)def predict(reading: Reading):    ...

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

Python
MAX_BYTES = 5 * 1024 * 1024@app.post("/predict/image", response_model=Prediction)def predict_image(file: UploadFile = File(...)):    if file.content_type not in {"image/jpeg", "image/png"}:        raise HTTPException(415, f"unsupported content type: {file.content_type}")    raw = file.file.read(MAX_BYTES + 1)    if len(raw) > MAX_BYTES:        raise HTTPException(413, "image exceeds 5 MB")    try:        img = Image.open(io.BytesIO(raw)).convert("RGB")    except Exception:        raise HTTPException(422, "file is not a decodable image")    x = STATE["tfm"](img).unsqueeze(0)    t0 = time.perf_counter()    with torch.no_grad():        logits = STATE["model"](x)        probs = torch.softmax(logits, dim=1)[0]    idx = int(probs.argmax())    return Prediction(label=STATE["labels"][idx],                      confidence=round(float(probs[idx]), 4),                      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.

DeclarationRuns onBlocking 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 loopCatastrophic — 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:

Python
import anyio, httpx@app.post("/predict/enriched")async def predict_enriched(req: Reading):    async with httpx.AsyncClient(timeout=2.0) as client:        r = await client.get(f"http://feature-store/user/{req.user_id}")   # real await    features = build(req, r.json())    probs = await anyio.to_thread.run_sync(run_model, features)            # off the loop    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:

ConfigurationSingle-request latencyThroughput 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:

Python
from fastapi.responses import JSONResponsefrom starlette.middleware.base import BaseHTTPMiddlewareimport uuid, logginglog = logging.getLogger("serving")@app.exception_handler(Exception)async def unhandled(request, exc):    rid = getattr(request.state, "request_id", "unknown")    log.exception("unhandled error request_id=%s path=%s", rid, request.url.path)    return JSONResponse(status_code=500,                        content={"error": "internal error", "request_id": rid})class Context(BaseHTTPMiddleware):    async def dispatch(self, request, call_next):        rid = request.headers.get("x-request-id") or str(uuid.uuid4())        request.state.request_id = rid        t0 = time.perf_counter()        response = await call_next(request)        ms = (time.perf_counter() - t0) * 1000        response.headers["x-request-id"] = rid        response.headers["x-process-time-ms"] = f"{ms:.1f}"        log.info("%s %s %d %.1fms rid=%s", request.method, request.url.path,                 response.status_code, ms, rid)        return responseapp.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:

Python
import timefrom collections import defaultdictfrom fastapi import HeaderRATE, BURST = 10.0, 20          # 10 requests/second sustained, bursts up to 20_buckets = defaultdict(lambda: [BURST, time.monotonic()])def check_rate(api_key: str = Header(..., alias="x-api-key")):    tokens, last = _buckets[api_key]    now = time.monotonic()    tokens = min(BURST, tokens + (now - last) * RATE)    if tokens < 1:        retry = (1 - tokens) / RATE        raise HTTPException(429, "rate limit exceeded",                            headers={"Retry-After": str(int(retry) + 1)})    _buckets[api_key] = [tokens - 1, now]    return api_key@app.post("/predict", dependencies=[Depends(check_rate)])def 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:

Python
@app.post("/predict")async def predict(reading: Reading):    try:        with anyio.fail_after(5.0):            return await anyio.to_thread.run_sync(                run_model, reading, abandon_on_cancel=True)    except TimeoutError:        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:

Python
from fastapi.responses import StreamingResponseimport json@app.post("/generate")def generate(req: GenRequest):    def token_stream():        for tok in model.stream(req.prompt, max_tokens=req.max_tokens):            yield f"data: {json.dumps({'token': tok})}\n\n"        yield "data: [DONE]\n\n"    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:

Python
from fastapi import BackgroundTasks@app.post("/predict")def predict(reading: Reading, bg: BackgroundTasks):    result = run_model(reading)    bg.add_task(record_prediction, reading.model_dump(), result)   # runs after response    return result

This 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:

Python
import threading_lock = threading.Lock()@app.post("/admin/reload")def reload_model(version: str, _=Depends(require_admin)):    new = torch.jit.load(f"/models/{version}/model.pt", map_location="cpu")    new.eval()    with torch.no_grad():                                   # prove it works first        new(torch.zeros(1, 3, 224, 224))    with _lock:        STATE["model"], STATE["version"] = new, version     # atomic swap    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

FlaskFastAPI
ModelWSGI, synchronousASGI, async-capable
ValidationHand-written or MarshmallowPydantic, from type hints
API docsAdd Flask-RESTX or write themOpenAPI + Swagger UI automatically
Concurrent I/O callsOne thread eachFree with await
Streaming responsesAwkward under WSGINative
Best fitSimple CPU-bound endpoint inside an existing Flask estateNew 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.