Model Deployment: From Notebook to Production
1 · The lesson
readRuntime note — FastAPI, TensorFlow Serving, Triton, and the rest of the deployment stack run in real Python/Docker environments. Pyodide won't help here. The snippets below are designed to be copied into a project, not pasted into a notebook.
A model in a notebook is worth nothing. A model behind an API with a 99.9% SLA, p95 latency under 200ms, autoscaling, monitoring, and a rollback button is worth its weight in revenue. The gap between the two is what this lesson is about.
dl-deploy covered the basics — load a model, write a Flask route. This lesson goes the rest of the way: serialisation formats, serving frameworks, inference optimisations, hosting platforms, cost math, batching, caching, monitoring, drift detection, and rollback. The stack you'll actually run in production.
1. The Deployment Maturity Model
Five stages, in roughly the order teams climb them:
| Stage | What it looks like | When it breaks |
|---|---|---|
| 1. Notebook | A Jupyter cell that predict()s. | The moment anyone else needs to call it. |
| 2. Script | A Python file someone SSHs into a box and runs. | Concurrency, restarts, monitoring. |
| 3. API | FastAPI/Flask service, one process, one machine. | Traffic exceeds one box. |
| 4. Containerised service | Docker image behind a load balancer; multiple replicas; orchestrated (Kubernetes, ECS). | Cost and operational burden grow. |
| 5. Managed serving | SageMaker / Vertex / Modal / Replicate. They handle the boring parts. | When you outgrow their pricing or constraints. |
Most teams should aim for stage 4 or 5. Skipping straight to Kubernetes for a model serving 10 RPS is over-engineering; staying at stage 2 for a 10 k RPS product is professional negligence.
2. Serialisation — Choose the Right Format
A trained model is weights plus architecture plus (sometimes) a preprocessing graph. The serialisation format determines what you can do with it.
| Format | Framework | Use when |
|---|---|---|
.keras (or .h5) | Keras | Single-framework Python serving with full custom-object support. |
state_dict (.pt / .pth) | PyTorch | Most flexible — weights only, you reconstruct the architecture in code. Standard for PyTorch serving. |
| TorchScript | PyTorch | Production. Captures the graph; runs without Python in C++ / mobile. |
| TensorFlow SavedModel | TF/Keras | Production. The native format for TF Serving. |
| ONNX | Cross-framework | You trained in PyTorch but want to serve in TensorRT, ONNX Runtime, CoreML, or any inference engine you didn't train in. |
| GGUF | llama.cpp et al. | LLM-specific quantised format for CPU and Apple Silicon inference. |
The rule of thumb: train in PyTorch/TF, serve in the engine that's fastest on your target hardware. ONNX is the lingua franca that bridges them. For non-LLM models, exporting to ONNX once and serving from ONNX Runtime usually gives you 2–5× speed-up vs serving from the training framework.
3. Serving Frameworks — Pick One
| Framework | Strengths | Notes |
|---|---|---|
| TF Serving | Mature, fast, gRPC + REST, version management, dynamic batching. | TensorFlow-only. |
| TorchServe | Same shape for PyTorch. | Less feature-complete than TF Serving. |
| NVIDIA Triton | Multi-framework (TF, PT, ONNX, TRT). Dynamic batching, model ensembles, GPU optimisations. | The serious choice for GPU-heavy production. Complex to operate. |
| BentoML | Pythonic. Build, package, deploy with a clear SDK. Multi-framework. | Popular with ML teams that don't want to learn Kubernetes day one. |
| LitServe | Lightning's. Modern, async, batches automatically, FastAPI under the hood. | Newer; rapidly growing. |
| vLLM / TGI | LLM-specific. Continuous batching, paged attention, optimised KV cache. | If you're serving an LLM, do not write your own. |
For a non-LLM model, BentoML or LitServe is the right starting point for a small ML team. Triton when you have GPU-saturating traffic or need the multi-framework / ensemble features. TF Serving / TorchServe when you're committed to one framework and want battle-tested infrastructure.
For LLMs, use vLLM or TGI (Text Generation Inference) — never roll your own. Continuous batching alone is a 5–10× throughput win that's prohibitive to reimplement.
4. Building a Real API with FastAPI
For a small model or as the front-end of a serving framework, FastAPI is the right starting point. The pattern:
# server.py from fastapi import FastAPI, HTTPException from pydantic import BaseModel, Field from contextlib import asynccontextmanager import numpy as np import time import logging import tensorflow as tf log = logging.getLogger("predict") logging.basicConfig(level=logging.INFO) class PredictRequest(BaseModel): # Pydantic does schema validation for free image: list[list[list[float]]] = Field(..., description="HxWx3 RGB pixels in [0,1]") class PredictResponse(BaseModel): label: str probability: float latency_ms: float # Load the model ONCE at startup — never per request state = {} @asynccontextmanager async def lifespan(app: FastAPI): log.info("loading model...") state["model"] = tf.keras.models.load_model("artifacts/cats_dogs.keras") state["labels"] = ["cat", "dog"] log.info("model loaded; ready") yield state.clear() app = FastAPI(lifespan=lifespan) @app.get("/health") def health(): return {"status": "ok", "model_loaded": "model" in state} @app.post("/predict", response_model=PredictResponse) def predict(req: PredictRequest): t0 = time.perf_counter() arr = np.asarray(req.image, dtype=np.float32) if arr.shape != (160, 160, 3): raise HTTPException(400, f"expected (160,160,3), got {arr.shape}") probs = state["model"].predict(arr[None, ...], verbose=0)[0] idx = int(np.argmax(probs)) latency_ms = (time.perf_counter() - t0) * 1000 log.info("predict", extra={"label": state["labels"][idx], "prob": float(probs[idx]), "latency_ms": latency_ms}) return PredictResponse( label=state["labels"][idx], probability=float(probs[idx]), latency_ms=latency_ms, )
Run with uvicorn server:app --workers 4 --port 8000. Five things to internalise:
lifespanloads the model at startup, not per request. Loading takes seconds; doing it per request is the single most common deployment mistake.- Pydantic validates the input schema. Wrong shape → 422 with a useful error. No model crashes on bad input.
- Health check so your load balancer can decide whether to route to this replica.
- Structured logging of the inputs, outputs, and latency. Without this you cannot debug production bugs.
--workers 4runs four independent processes — separate model copies, separate GPUs if you have them. CPU-bound inference scales with workers.
async def doesn't help for CPU-bound model inference — Python's GIL ensures only one prediction runs at a time per process. Use sync def (FastAPI runs it in a threadpool) and scale with --workers. async is the right answer for IO-heavy serving (calling other models, fetching from a DB) — combine the two where appropriate.
5. Inference Optimisations
A trained model is rarely deployment-ready as-is. Three families of optimisation, in increasing order of effort:
Quantisation (FP32 → INT8)
Convert weights from 32-bit floats to 8-bit integers. Roughly:
- 4× smaller on disk and in memory.
- 2–4× faster inference, often more on CPU.
- Small accuracy loss — usually under 1% if you do post-training quantisation, often zero if you do quantisation-aware training.
For PyTorch:
import torch from torch.quantization import quantize_dynamic quantised = quantize_dynamic(model, {torch.nn.Linear}, dtype=torch.qint8) torch.save(quantised.state_dict(), "model_int8.pt")
setup added so this can run · defines model
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) model = _AutoMock('model')
For ONNX, use onnxruntime.quantization.quantize_dynamic. For LLMs, formats like GPTQ, AWQ, and bitsandbytes' NF4 are the modern alternatives — same idea, more careful per-layer calibration.
Pruning and Distillation
- Pruning zeros out small weights. Combined with a sparse-aware runtime, 50–90% sparsity is achievable for modest accuracy cost. Often skipped in favour of quantisation, which is easier.
- Distillation trains a smaller "student" model to mimic a larger "teacher". DistilBERT is BERT distilled — 40% smaller, 60% faster, 97% of the accuracy. Worth the trouble when you'll serve the model billions of times.
Compilation
torch.compile(model)— PyTorch 2.x. Traces the model, optimises with Inductor, can give 30–80% speed-ups with one line.- TF
@tf.function— graph compilation; default for SavedModel inference. - ONNX Runtime — load an ONNX model, get hardware-specific acceleration (CPU, CUDA, TensorRT) for free.
- TensorRT — NVIDIA-specific. The fastest GPU inference. Significant export effort; reserve for high-volume production.
Stack these. Quantised + compiled + served from ONNX Runtime is a common production configuration and routinely 10× faster than the same model running model.predict() in Python.
6. Hosting Platforms 2026
| Platform | Model | Best for |
|---|---|---|
| SageMaker / Vertex AI | Managed, "bring your own container" | Teams already in AWS/GCP; complex MLOps integrations; willing to pay for managed everything. |
| Modal | Serverless GPU, Python-native | Bursty workloads, prototypes, batch jobs. Cold starts measured in seconds. |
| Replicate | Model marketplace + serverless | Diffusion / LLM demos. Per-second GPU billing. |
| Hugging Face Inference Endpoints | Managed HF models | Anything you'd run from the HF hub. |
| RunPod / Lambda Labs | Raw GPU rental | When you want full control and the lowest GPU-hour price. |
| Self-hosted on EC2 / GCE / Hetzner | DIY | When you have ops capacity, predictable load, and want to control costs. |
The honest sequencing for most teams:
1. Start on Modal or HF Inference Endpoints — zero ops, ship in a day.
2. When monthly bills exceed the cost of one engineer-week, migrate to SageMaker / Vertex for managed scale.
3. When monthly bills exceed the cost of one engineer-month, build out self-hosted infrastructure on EC2/GCE with proper autoscaling.
Premature self-hosting is a classic small-team mistake.
7. Cost — The "Always-On GPU" Question
A simple model: a T4 on AWS is roughly $0.50/hour. That's $360/month if you run it 24/7, regardless of traffic. If your service averages 5 RPS and the model takes 50ms per request, you're using 25% of the GPU and paying for the other 75%.
The serverless alternative (Modal, Replicate, Banana) bills per second of GPU use. The same workload costs maybe $80/month. Cold starts are 2–10 seconds — fine for batch, painful for interactive.
Rough rules:
- Steady traffic > 30% GPU utilisation → always-on GPU is cheaper.
- Spiky traffic, latency-tolerant → serverless GPU.
- Latency-critical, low-volume → always-on, dedicated.
- Latency-critical, high-volume → always-on with autoscaling.
Always-on with autoscaling means: keep a minimum of N replicas warm, scale up beyond that under load. The minimum kills your cold-start problem; the scaling caps your bill.
Per-1000-requests math is the right unit to argue about budgets in. Take total monthly cost, divide by monthly requests, multiply by 1000. Compare to competitors and SaaS alternatives. If your model costs more per 1000 requests than buying the same inference from OpenAI's API, you have an architecture problem.
8. Batching at Inference
GPUs are designed to do the same operation on many inputs at once. A batch of 32 takes barely longer than a batch of 1. Inference servers exploit this with dynamic batching: hold incoming requests for a few milliseconds, group them, run one forward pass, return all results.
Triton, TF Serving, LitServe, and BentoML all support dynamic batching out of the box. Two knobs:
- Maximum batch size — caps memory.
- Maximum wait time — caps latency. Common values: 5–20ms. Higher = better throughput, worse p99.
A model that took 50ms per request unbatched can serve 10× the QPS at the cost of an extra 10ms of latency. This is the most lopsided optimisation you'll do — and you'll see it leave a 10× speed-up on the table on any service that doesn't batch.
9. Caching
For deterministic models with cacheable inputs, caching identical-input requests is free QPS. The implementation is straightforward:
import hashlib from functools import lru_cache def cache_key(image_array: np.ndarray) -> str: return hashlib.sha256(image_array.tobytes()).hexdigest() cache: dict[str, dict] = {} def predict_cached(image_array): key = cache_key(image_array) if key in cache: return cache[key] | {"cache": "hit"} result = run_model(image_array) cache[key] = result return result | {"cache": "miss"}
setup added so this can run · defines run_model, np
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) def run_model(*_a, **_kw): print('-> run_model() called') return _AutoMock('run_model()') np = _AutoMock('np')
In production replace the in-memory dict with Redis (multi-replica) and add an LRU/TTL eviction policy. For high-cardinality inputs (images, text > a few hundred characters), the cache hit rate is low and the bookkeeping overhead might cost more than it saves. For low-cardinality inputs (canonical product IDs, repeated queries), cache hit rates can exceed 90%.
10. Monitoring
Three layers, all required:
Latency Histograms
p50, p95, p99 of request latency. Mean is misleading — a few slow requests can ruin user experience while the mean stays low.
from prometheus_client import Histogram, start_http_server LATENCY = Histogram("predict_latency_seconds", "Prediction latency") @LATENCY.time() def predict(...): ...
Input Distribution Drift
Track summary statistics of incoming inputs — mean, std, the rate of out-of-range values, the distribution of categorical features. Compare to training-data statistics. If your image-classifier suddenly starts seeing inputs with mean brightness 30% lower than training, something upstream broke (camera change, codec change, edge-cropping bug) and predictions will silently degrade.
This is covariate shift. The model isn't wrong; the world it operates in has changed.
Output Distribution Drift
The class balance of your predictions, the distribution of regression outputs, the rate of confident vs uncertain calls. A spam classifier that suddenly classifies 80% of email as spam (up from 5%) didn't get smarter — something is broken.
This is concept drift (the relationship between inputs and outputs changed) or a feedback loop (your earlier predictions changed user behaviour, which changed the inputs).
Shadow Evaluation
Where possible, run a slice of live traffic against a hold-out labelled set or against a human-labelled "shadow" set and report ongoing model quality. This is the only monitoring that catches genuine quality regressions.
11. Rollback Strategies
When a new model version misbehaves in production, you want it gone in seconds, not hours.
- Blue/green deployment — keep the old version running, route 100% to new. Switch back instantly if needed.
- Canary deployment — route 1% / 5% / 25% of traffic to the new version, watch metrics, expand or roll back. The default for serious ML deployment.
- A/B testing — run two models concurrently on different cohorts, measure business metrics, pick a winner. Different goal than canary (statistical comparison vs safety), often run as the same machinery.
- Shadow mode — run the new model on every request alongside the old one, but only serve the old one's predictions. Use the captured predictions to measure new-model behaviour at zero user risk.
The combination most teams settle into: shadow mode for an offline validation period, then a canary rollout to 5% → 25% → 100% over a few days, with automatic rollback if any monitored metric breaches a threshold.
Common Mistakes
1. Loading the model on every request.
You write model = load_model(path) inside the route handler "for clarity". Latency goes from 50ms to 5 seconds because each request reloads several hundred megabytes from disk. Use lifespan (FastAPI) or @app.on_event("startup") (older FastAPI) or module-level loading in a singleton — never per request.
2. No batching.
A model that could serve 1000 RPS in batches of 32 instead serves 100 RPS one at a time. Add dynamic batching at the serving layer; the throughput multiplier is almost always 5–20×.
3. Single point of failure.
One replica behind one load balancer. The replica restarts, the load balancer drops requests, the on-call gets paged at 3 AM. Run at least two replicas, ideally in different availability zones, behind a load balancer that supports automatic unhealthy-instance removal.
4. No input validation.
The route accepts arbitrary JSON, passes it to the model, and crashes when someone POSTs a dict where you expected a list. Worst case it doesn't crash — it returns confident nonsense for malformed input. Use Pydantic / JSON Schema to validate at the boundary.
5. Unpinned library versions.
tensorflow>=2.10 in requirements.txt looked fine when 2.11 was released. Then 2.15 changed a default and your serialised model loads but produces different outputs. Pin exact versions (tensorflow==2.15.0) in your deployed image. Reproducibility beats convenience.
6. Logging predictions but not inputs.
You log "predicted label=cat, prob=0.92" but not the image. Three days later a customer complains "your model said my dog was a cat" and you can't reproduce the bug because the input is gone. Log inputs too — to object storage if they're large, with retention policies for privacy. You cannot debug ML in production without the inputs.
7. Treating "model accuracy on the holdout set" as production quality.
Holdout accuracy is a floor on production quality, not a measurement of it. Real production sees out-of-distribution inputs, adversarial inputs, the long tail of edge cases your test set didn't represent. Shadow evaluation, drift monitoring, and labelled live samples are how you measure real quality.
🎯 Your Turn — A Production-Shaped FastAPI Endpoint
Write a FastAPI endpoint that serves a Keras image classifier. Requirements:
- Loads the model once at startup using FastAPI's
lifespancontext manager. - Accepts a JSON request with a 160×160×3 list-of-lists image (Pydantic-validated).
- Validates that the input has the right shape; returns a 400 if not.
- Returns a JSON response containing the predicted label, the probability, and the latency in milliseconds.
- Includes a
GET /healthendpoint that returns{"status": "ok"}once the model is loaded.
Skeleton:
# server.py from fastapi import FastAPI, HTTPException from pydantic import BaseModel, Field from contextlib import asynccontextmanager import numpy as np import time import tensorflow as tf LABELS = ["cat", "dog"] INPUT_SHAPE = (160, 160, 3) # TODO 1: Pydantic request and response models class PredictRequest(BaseModel): ... class PredictResponse(BaseModel): ... # TODO 2: lifespan context manager — load model into a state dict state = {} @asynccontextmanager async def lifespan(app: FastAPI): ... app = FastAPI(lifespan=lifespan) # TODO 3: /health endpoint # TODO 4: /predict endpoint — validate shape, run model, time it, return response
Hint 1 — Lifespan and module-level state
The model goes into a module-level dict (or a single attribute onapp.state). Inside lifespan: load before the yield, clear after. state["model"] = tf.keras.models.load_model("..."). The handler then reads state["model"].
Hint 2 — Timing with perf_counter
t0 = time.perf_counter() at the top of the handler, (time.perf_counter() - t0) * 1000 at the end. Always use perf_counter for measuring durations — it's monotonic and high-resolution.
Show full solution
# server.py from fastapi import FastAPI, HTTPException from pydantic import BaseModel, Field from contextlib import asynccontextmanager import numpy as np import time import logging import tensorflow as tf LABELS = ["cat", "dog"] INPUT_SHAPE = (160, 160, 3) log = logging.getLogger("predict") logging.basicConfig(level=logging.INFO) class PredictRequest(BaseModel): image: list[list[list[float]]] = Field( ..., description="160x160x3 RGB image in [0, 1]" ) class PredictResponse(BaseModel): label: str probability: float latency_ms: float state: dict = {} @asynccontextmanager async def lifespan(app: FastAPI): log.info("loading model...") state["model"] = tf.keras.models.load_model("artifacts/cats_dogs.keras") state["ready"] = True log.info("model loaded; ready") yield state.clear() app = FastAPI(lifespan=lifespan) @app.get("/health") def health(): return {"status": "ok" if state.get("ready") else "loading"} @app.post("/predict", response_model=PredictResponse) def predict(req: PredictRequest): t0 = time.perf_counter() arr = np.asarray(req.image, dtype=np.float32) if arr.shape != INPUT_SHAPE: raise HTTPException( status_code=400, detail=f"expected shape {INPUT_SHAPE}, got {arr.shape}", ) probs = state["model"].predict(arr[None, ...], verbose=0)[0] if probs.shape == (1,): # binary sigmoid p_dog = float(probs[0]) idx = 1 if p_dog >= 0.5 else 0 prob = p_dog if idx == 1 else 1 - p_dog else: # softmax over classes idx = int(np.argmax(probs)) prob = float(probs[idx]) latency_ms = (time.perf_counter() - t0) * 1000 log.info( "predict", extra={"label": LABELS[idx], "prob": prob, "latency_ms": latency_ms}, ) return PredictResponse( label=LABELS[idx], probability=prob, latency_ms=latency_ms, )
Run it:
uvicorn server:app --workers 4 --port 8000
curl http://localhost:8000/health
# {"status":"ok"}
curl -X POST http://localhost:8000/predict \
-H "Content-Type: application/json" \
-d "$(python -c 'import json,numpy as np; print(json.dumps({"image": np.zeros((160,160,3)).tolist()}))')"
# {"label":"cat","probability":0.512,"latency_ms":47.31}This is the shape of a production service. To make it actually production-ready, layer on the rest of the lesson: structured logging to stdout, Prometheus metrics, a Dockerfile, a /ready probe distinct from /health, a docker-compose with a load balancer in front, dynamic batching at the serving layer (or migrate to BentoML/LitServe), and canary deployment via your orchestrator. The skeleton in this exercise is the foundation everything else sits on.
What You Learned
- The deployment maturity model: notebook → script → API → containerised service → managed serving. Most teams should aim for stage 4 or 5.
- Pick the right serialisation format for the target runtime —
state_dictfor PyTorch dev, TorchScript / SavedModel for serving, ONNX to cross frameworks, GGUF for LLM CPU inference. - Serving frameworks: TF Serving, TorchServe, Triton (heavyweight), BentoML / LitServe (Pythonic), vLLM / TGI (LLM-specific). Don't write your own batching loop.
- A FastAPI service is the right starting point for small models. Load the model at startup, validate inputs with Pydantic, return structured responses, log everything.
- Inference optimisations stack: quantisation (FP32 → INT8 ≈ 4× smaller, 2–4× faster), distillation,
torch.compile/tf.function, ONNX Runtime, TensorRT. - Hosting platforms 2026: Modal / Replicate / HF Endpoints for ease, SageMaker / Vertex for managed scale, self-hosted EC2 for cost control. Premature self-hosting kills small teams.
- Cost-per-1000-requests is the right unit. Always-on GPU vs serverless GPU is a utilisation question.
- Dynamic batching is the most lopsided optimisation in serving — 5–20× throughput for ~10ms of latency.
- Monitoring: latency histograms (p50/p95/p99), input drift, output drift, shadow evaluation. Mean latency lies.
- Rollback strategies: blue/green, canary, A/B, shadow mode. Real ML deployment is canary by default.
Next: AI Ethics: Building Responsibly — the harms ML systems cause at scale, how to measure them, and the design choices that prevent them.