Skip to content

Online Serving with FastAPI + Docker

At a glance A complete, hands-on guide to wrapping a trained PyTorch model in a FastAPI service and shipping it as a Docker image: model export, pre/post-processing wrappers, gunicorn multi-process deployment, health checks, and the usual pitfalls.

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")
EOF

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

  1. 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;
  2. 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;
  3. torch.inference_mode() is faster than torch.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: str
python
# 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 /predict takes one image by default, because an online endpoint has to protect latency; batch inference lives in a dedicated POST /predict_batch that 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_k and predictions — stable fields with clear semantics, so callers never have to parse a black-box JSON blob.
  • Health checks split into livez and health: 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 locationCost per requestBehavior
Inside the request handler2-4s + inferenceEvery request is slow enough to time out; GPU/memory usage fluctuates
Module level0 (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/
*.pth

Should 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: 60s

7. 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.py
python
# 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: 10

Don'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/predict

For 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 ​

PitfallSymptomFix
Model loaded inside the request handlerEvery request 2-4s slower; memory keeps climbingMove it to module level — see Section 5
Worker count × model GPU memory > total GPU memoryRandom CUDA out of memory after startupCount the workers, or probe GPU memory before startup
Synchronous inference inside async defLatency degrades linearly as concurrency risesSwitch to def or a thread pool — see the Section 4 warning
gunicorn's default timeout=30sSlow requests return 502Set timeout to 3-5x the P99 latency
Missing model.eval()Accuracy drops mysteriously in productionCall eval explicitly before inference
Weights baked into the imageImage is several GB; CI is slowExclude via .dockerignore + mount at runtime
Multi-worker cold startPod restarts over and overAdd a startupProbe — see Section 8

Further Reading ​

References ​