Appearance
Online Serving with FastAPI + Docker: From a Trained PyTorch Model to a Production API
In one sentence: use FastAPI to wrap a trained PyTorch model into an HTTP inference service, then package it into a shippable, scalable Docker image — this is the shortest path from "the model runs on my laptop" to "the model reliably serves business traffic."
Why it's worth doing: FastAPI + Docker is the most mainstream combination for small and mid-size teams putting models into production, covering 80% of the fundamentals of serving — model loading, the API contract, the process model, image building, health checks. Once you've run this pipeline end to end, your understanding of Serving and Inference APIs shifts from "I've heard of it" to "I've built it"; when the models multiply and throughput falls short, you can upgrade smoothly to dedicated inference servers such as NVIDIA Triton or vLLM. This article uses taking an image classification model to production as its scenario, and the whole pipeline is reproducible.
1. The Scenario: Shipping an Image Classification Service
Suppose we're launching a ResNet18 image classification service: a client (a mobile app) uploads an image, and the service returns the Top-5 classes with confidence scores. The constraints:
- The model runs inference on CPU or a single GPU; the per-request P99 latency target is < 100ms;
- Production traffic is moderate QPS (tens to hundreds), with occasional bursts;
- The deliverable is a Docker image that runs on the company's container platform (K8s or Compose).
The plan up front: the whole project breaks into five pieces — model export, the predict wrapper class, the FastAPI app, the gunicorn process model, and the Docker image. We'll build them in that order.
2. Preparing the Model: torchvision Pretrained Weights
No real training needed — take torchvision's ResNet18 pretrained weights and export them into a standalone file:
bash
python - <<'EOF'
import torch
from torchvision.models import resnet18, ResNet18_Weights
model = resnet18(weights=ResNet18_Weights.DEFAULT) # ImageNet1K pretrained
model.eval()
torch.save(model.state_dict(), "resnet18.pth")
print("saved", sum(p.numel() for p in model.parameters()) / 1e6, "M params")
EOFThe key point: save the state_dict, not the whole model object. Pickling the whole object serializes the class definitions and the torchvision version along with the weights, and it frequently fails to deserialize after a dependency upgrade; a state_dict is just a dictionary of tensors and is portable across environments. This is the same principle as "freeing the model from its training environment" in Model Formats and Conversion — except that here we've only achieved "portable weights, still tied to PyTorch"; exporting to ONNX takes it further when needed.
3. Exporting and Wrapping the Model: the predict Class
Wrapping "model + preprocessing + postprocessing" into a single class is one of the most important habits in serving engineering. The reason: a stable interface contract. You can swap out pre/post-processing without changing the service's external schema; callers only ever interact with predict(image_bytes) -> [{"label","score"}] and never touch tensors.
python
# app/model.py
import io
import torch
import torchvision.transforms as T
from PIL import Image
from torchvision.models import resnet18
# The 1000 ImageNet-1K class names (load from a file in real projects)
CLASS_NAMES = [f"class_{i}" for i in range(1000)]
class ImageClassifier:
"""Bundles model, preprocessing, and postprocessing into one class; only predict is exposed"""
def __init__(self, model_path: str, device: str = "cpu"):
self.device = torch.device(device)
# weights=None: don't load torchvision's bundled weights; load our exported state_dict instead
self.model = resnet18(weights=None)
self.model.load_state_dict(
torch.load(model_path, map_location=self.device)
)
self.model.to(self.device)
self.model.eval() # disable dropout / BN running-stats updates
# Preprocessing must match training exactly: Resize(256) -> CenterCrop(224) -> Normalize
self.transform = T.Compose([
T.Resize(256),
T.CenterCrop(224),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
def preprocess(self, image_bytes: bytes) -> torch.Tensor:
img = Image.open(io.BytesIO(image_bytes)).convert("RGB")
tensor = self.transform(img).unsqueeze(0) # [1, 3, 224, 224]
return tensor.to(self.device)
def postprocess(self, logits: torch.Tensor, top_k: int = 5) -> list[dict]:
probs = torch.softmax(logits, dim=1)[0]
topk = torch.topk(probs, k=top_k)
return [
{"label": CLASS_NAMES[int(i)], "score": round(float(s), 4)}
for s, i in zip(topk.values, topk.indices)
]
@torch.inference_mode() # skips gradient tracking, lowering GPU/memory usage
def predict(self, image_bytes: bytes, top_k: int = 5) -> list[dict]:
tensor = self.preprocess(image_bytes)
logits = self.model(tensor)
return self.postprocess(logits, top_k)Three easy-to-get-wrong spots, flagged up front:
model.eval()must be called explicitly — skip it and BatchNorm uses per-batch statistics instead of running statistics at inference time, which can cost a few points of accuracy;- Every preprocessing parameter must line up with training (the Resize size, the Normalize mean/std) — this is the number-one source of the "training-serving skew" listed in Common Pitfalls and Anti-Patterns;
torch.inference_mode()is faster thantorch.no_grad()because it doesn't even build the autograd metadata inference needs.
4. API Design with FastAPI
Define a clear input/output schema (with Pydantic), then write the endpoints. Contract first, logic second:
python
# app/schemas.py
from pydantic import BaseModel
class Prediction(BaseModel):
label: str
score: float
class PredictResponse(BaseModel):
top_k: int
predictions: list[Prediction]
class HealthResponse(BaseModel):
status: str
model_loaded: bool
device: strpython
# app/main.py
from fastapi import FastAPI, File, HTTPException, UploadFile
from app.model import ImageClassifier
from app.schemas import HealthResponse, PredictResponse
# Module-level load: loaded once at process startup; all requests reuse the same weights
classifier = ImageClassifier(model_path="/models/resnet18.pth", device="cuda:0")
app = FastAPI(title="image-classifier", version="1.0.0")
@app.post("/predict", response_model=PredictResponse)
async def predict(file: UploadFile = File(...), top_k: int = 5) -> PredictResponse:
"""Upload an image, get the Top-K classification back"""
data = await file.read()
if not data:
raise HTTPException(status_code=400, detail="empty image")
predictions = classifier.predict(data, top_k=top_k)
return PredictResponse(top_k=top_k, predictions=predictions)
@app.get("/health", response_model=HealthResponse)
def health() -> HealthResponse:
"""Readiness probe: ready only when the model is loaded and the device works"""
return HealthResponse(status="ok", model_loaded=True, device=str(classifier.device))
@app.get("/livez")
def livez() -> dict:
"""Liveness probe: 200 as long as the process is alive; no heavyweight checks"""
return {"status": "alive"}Design decisions and rationale:
- Single-item first; batching gets its own endpoint.
POST /predicttakes one image by default, because an online endpoint has to protect latency; batch inference lives in a dedicatedPOST /predict_batchthat assembles a batch internally and runs a single forward pass. This matches the "single vs batch" discussion in Serving and Inference APIs: online traffic is naturally scattered, so it's worth accumulating server-side and computing in bulk. - The response carries
top_kandpredictions— stable fields with clear semantics, so callers never have to parse a black-box JSON blob. - Health checks split into
livezandhealth: the former only proves the process is alive, the latter that the model is ready. They map onto K8s liveness and readiness probes respectively, preventing the case where "the process is up, the model is still loading, and traffic arrives anyway."
The synchronous blocking problem
predict does CPU/GPU-heavy synchronous inference. If the endpoint is declared async def predict, inference will block the entire event loop — with 10 concurrent requests, the other 9 all queue up. Two ways out: declare the endpoint as a plain def (FastAPI automatically dispatches it to a thread pool), or use run_in_executor. At high concurrency the more thorough fix is to move to an engine with native concurrent batching such as Triton. The example above keeps async def only to demonstrate asynchronous file reading; refactor it this way before any real load testing.
5. Where the Model Loads: Why It Must Be Module-Level
Bottom line: load the model once at process startup (module level) — never inside the request handler. Putting ImageClassifier(...) at the top of the module executes it once when the process starts/forks; every worker and every request shares that one in-memory copy.
Compare the two placements (ResNet18 has roughly 45M parameters; weight loading + state_dict deserialization takes about 2-4 seconds):
| Load location | Cost per request | Behavior |
|---|---|---|
| Inside the request handler | 2-4s + inference | Every request is slow enough to time out; GPU/memory usage fluctuates |
| Module level | 0 (pay only the inference cost) | Stable, constant memory footprint |
There's one more engineering reason for module-level loading: load failures should surface at startup, not as a 500 on the first request. Combined with the container platform's startupProbe, an image that isn't ready within tens of seconds gets killed and restarted, so the problem surfaces during deployment instead of exploding when traffic arrives.
6. Dockerizing: Multi-Stage Builds, Non-Root, .dockerignore
1. Dockerfile (multi-stage build)
dockerfile
# ===== Stage 1: builder — installs dependencies only =====
FROM python:3.11-slim AS builder
WORKDIR /app
COPY requirements.txt .
# Install the CPU build of torch from the official index: ~1GB smaller than the default CUDA build on PyPI
RUN pip install --no-cache-dir --index-url https://download.pytorch.org/whl/cpu \
torch==2.3.0 torchvision==0.18.0 && \
pip install --no-cache-dir -r requirements.txt
# ===== Stage 2: runtime — the minimal running image =====
FROM python:3.11-slim AS runtime
ENV PYTHONUNBUFFERED=1 PIP_NO_CACHE_DIR=1
# Non-root: run as an unprivileged user inside the container to reduce security risk
RUN useradd -m -u 1000 appuser
WORKDIR /app
COPY --from=builder /usr/local/lib/python3.11/site-packages /usr/local/lib/python3.11/site-packages
COPY app/ app/
COPY models/resnet18.pth models/resnet18.pth
USER appuser
EXPOSE 8000
CMD ["gunicorn", "app.main:app", "-k", "uvicorn.workers.UvicornWorker", "-c", "gunicorn.conf.py"]2. .dockerignore
text
__pycache__/
*.pyc
.git/
.venv/
tests/
data/raw/
*.pthShould the weights go into the image?
The example above copies resnet18.pth into the image for convenience. The more common production pattern is an image without weights, pulled at runtime from object storage or a model registry, so weight updates don't require rebuilding the image. A middle ground is a multi-stage image with the weights in their own layer. Note that ignoring *.pth in .dockerignore implies the weights are externally mounted — don't do both at once and trip over each other.
3. docker-compose (with GPU and health check)
yaml
services:
inference:
build: .
ports:
- "8000:8000"
environment:
WEB_CONCURRENCY: "2"
TIMEOUT: "60"
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: 1
capabilities: [gpu]
healthcheck:
test: ["CMD", "python", "-c", "import urllib.request;urllib.request.urlopen('http://localhost:8000/health')"]
interval: 10s
timeout: 5s
retries: 3
start_period: 60s7. gunicorn + uvicorn: How Many Workers
uvicorn, the ASGI server behind FastAPI, defaults to a single process with one worker. A single worker can only saturate one core on CPU, so production almost always launches multiple uvicorn workers under gunicorn:
bash
gunicorn app.main:app -k uvicorn.workers.UvicornWorker -c gunicorn.conf.pypython
# gunicorn.conf.py
import os
workers = int(os.environ.get("WEB_CONCURRENCY", "2"))
worker_class = "uvicorn.workers.UvicornWorker"
bind = "0.0.0.0:8000"
timeout = int(os.environ.get("TIMEOUT", "60")) # the 30s default is too short; slow inference gets killed
graceful_timeout = 30
max_requests = 5000 # gracefully restart after 5000 requests to guard against memory leaks
max_requests_jitter = 500 # add jitter so all workers don't restart at the same time
accesslog = "-"
errorlog = "-"The worker-count trade-off, conclusions first:
- CPU deployment:
workers ≈ CPU core count(the model's memory footprint is small — one copy per worker, typically under 2GB — so you have plenty of headroom); - GPU deployment: workers = GPU memory budget ÷ per-worker memory. Each worker loads its own copy of the model and reserves its own CUDA context (a fixed overhead of roughly 300-500MB). ResNet18 FP32 weights are about 172MB, so per-worker GPU memory is around 600MB; a 24GB A10 could theoretically fit a dozen copies, but past 3-4 workers the returns diminish, because each infers independently and they never share a batch.
Rules of thumb: for a single-model GPU service, use 1-2 workers, then add in-process batching or move straight to Triton; on multi-core CPU machines stay at or below 2x the core count. The real number comes from measuring with the methods in Load Testing and Capacity Planning: add workers step by step and watch for the point where throughput stops climbing. Reference measurements: ResNet18 on an 8-core CPU machine — 1 worker ≈ 18 QPS, 2 workers ≈ 33 QPS, 4 workers ≈ 55 QPS (latency rises accordingly, 50ms → 90ms); on an A10 GPU a single worker already reaches 200+ QPS.
8. Health Checks and Probes (the Kubernetes Version)
yaml
containers:
- name: inference
image: registry.example.com/image-classifier:1.0.0
ports: [{ containerPort: 8000 }]
startupProbe: # model loading is slow; allow plenty of startup time
httpGet: { path: /health, port: 8000 }
failureThreshold: 60
periodSeconds: 5 # up to 300s total
livenessProbe: # restart when the process dies
httpGet: { path: /livez, port: 8000 }
initialDelaySeconds: 30
readinessProbe: # receive traffic only when ready
httpGet: { path: /health, port: 8000 }
periodSeconds: 10Don't skip startupProbe
With multiple workers, total cold-start time ≈ worker count × per-worker load time. Two workers each taking 4 seconds to load, plus dependency initialization, can leave the Pod not ready until 30-60 seconds in. Without a startupProbe, the readinessProbe starts probing on a 10s cycle from the very beginning, K8s restarts the Pod over and over, and you get a "never comes up" livelock. This is one of the most common launch failures for newcomers.
9. Load Testing and Capacity Validation
Run at least one round of load testing before launch to answer three questions: what's the P99 latency, what's the maximum throughput, and is the bottleneck the CPU or the GPU memory. Recommended tools: locust (Python ecosystem) or wrk (quick and simple):
bash
# wrk: 200 concurrent connections, 30-second run, script sends a real image
wrk -t 8 -c 200 -d 30s -s post.lua http://localhost:8000/predictFor the full method and metric interpretation (QPS, P99, error rate, GPU utilization), see Load Testing and Capacity Planning. While testing, watch GPU memory and utilization with nvidia-smi dmon — if GPU utilization stays under 30%, the bottleneck is CPU preprocessing or too few gunicorn workers, not the GPU.
Common Pitfalls and Troubleshooting
| Pitfall | Symptom | Fix |
|---|---|---|
| Model loaded inside the request handler | Every request 2-4s slower; memory keeps climbing | Move it to module level — see Section 5 |
| Worker count × model GPU memory > total GPU memory | Random CUDA out of memory after startup | Count the workers, or probe GPU memory before startup |
Synchronous inference inside async def | Latency degrades linearly as concurrency rises | Switch to def or a thread pool — see the Section 4 warning |
| gunicorn's default timeout=30s | Slow requests return 502 | Set timeout to 3-5x the P99 latency |
Missing model.eval() | Accuracy drops mysteriously in production | Call eval explicitly before inference |
| Weights baked into the image | Image is several GB; CI is slow | Exclude via .dockerignore + mount at runtime |
| Multi-worker cold start | Pod restarts over and over | Add a startupProbe — see Section 8 |
Further Reading
- Serving and Inference APIs — the theory behind API design, single vs batch, and protocol choice; this article is its worked example
- Load Testing and Capacity Planning — close the loop on this article's worker counts and QPS targets with real load-test data
- Model Formats and Conversion — the starting point for the next stage of optimization: PyTorch → ONNX → TensorRT
- Common Pitfalls and Anti-Patterns — the complete serving pitfall list; this article only covers the high-frequency items
- Multi-Model Serving with NVIDIA Triton — the upgrade path when a single-model FastAPI service can't handle the throughput or the model count
- Anatomy of an Inference System — placing this article's single service back into a full inference system to see where it fits
References
- FastAPI documentation: https://fastapi.tiangolo.com/
- Uvicorn documentation: https://www.uvicorn.org/
- Gunicorn documentation: https://docs.gunicorn.org/en/stable/
- PyTorch documentation: https://pytorch.org/docs/stable/
- Docker documentation: https://docs.docker.com/