Course Content
Model Deployment for AI Engineers
4 sections · 10 lessons
Deploy a CNN Model as a REST API
Here is the situation this project puts you in. A quality team at a small manufacturer photographs machined parts on a conveyor and wants a service that classifies each photo. They will call it from a Raspberry Pi on the line, over the local network, at about two images per second during a shift. They need an answer in under 200 milliseconds, they cannot install Python on the Pi beyond requests, and the person maintaining it after you leave has never used Docker.
That set of constraints — modest traffic, a hard latency budget, a non-expert operator, and a plain HTTP interface — describes an enormous share of real deployed machine learning. It is also small enough to build completely in a few weekends, which is why it is the right thing to build end to end at least once.
You will train a convolutional network on CIFAR-10 as the stand-in for the parts dataset, export it to a portable artefact, serve it behind FastAPI, containerise it, test it, ship it through CI, and deploy it somewhere reachable. The point is not the model. Any competent CIFAR-10 CNN will do. The point is everything that surrounds it.
What you are building, precisely
Write these numbers down before you start, because "it works" is not a finish line and without a target you will keep polishing.
| Requirement | Target | How you will measure it |
|---|---|---|
| Test accuracy | ≥ 85% top-1 on the CIFAR-10 test split | scripts/evaluate.py, held-out 10,000 images |
| Per-class accuracy | No class below 75% | Confusion matrix, not aggregate accuracy |
| Latency | p95 under 100 ms, batch size 1, 2 CPU cores | 200 timed requests after a warm-up |
| Container image | Under 1.5 GB | docker images |
| Cold start | Ready within 30 s of container start | Time from docker run to /readyz returning 200 |
| Robustness | Malformed input returns 4xx, never 500 | Integration tests with deliberate garbage |
| Observability | Every request has an ID and a structured log line | Read the logs after a load test |
Write your acceptance numbers before you write any code; a project without a definition of done becomes a project without an end.
Project layout
cifar-api/├── app/│ ├── __init__.py│ ├── main.py FastAPI application and routes│ ├── model.py load, preprocess, predict, postprocess│ ├── schemas.py Pydantic request/response models│ └── logging_conf.py JSON formatter, correlation IDs├── training/│ ├── train.py training loop, checkpointing│ ├── export.py TorchScript + ONNX export and verification│ └── evaluate.py accuracy, per-class metrics, confusion matrix├── models/│ ├── model.pt TorchScript artefact│ ├── model.onnx ONNX artefact│ └── metadata.json version, digest, metrics, opset, git commit├── tests/│ ├── unit/ preprocessing, schemas — model mocked│ ├── integration/ real container over HTTP│ ├── fixtures/ a dozen real images with known labels│ └── golden.npz inputs + expected outputs, byte-stable├── .github/workflows/ci.yml├── Dockerfile├── docker-compose.yml├── requirements.txt└── README.mdThe separation that matters most is app/model.py containing no FastAPI imports at all. Keep loading, preprocessing, inference, and postprocessing as plain functions, and the web layer becomes a thin adapter. That is what lets you later wrap the same core in a Lambda handler, a Gradio demo, or a batch script without touching the logic — and it is what makes the core unit-testable without spinning up a server.
Phase 1 — Train something good enough
Spend the least time here that gets you past 85%. A small VGG-style network trained for 30 epochs with augmentation comfortably reaches 87–90% on CIFAR-10.
1import torch, torch.nn as nn, torchvision2import torchvision.transforms as T34MEAN, STD = (0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)56train_tfm = T.Compose([7 T.RandomCrop(32, padding=4),8 T.RandomHorizontalFlip(),9 T.ToTensor(),10 T.Normalize(MEAN, STD),11])12test_tfm = T.Compose([T.ToTensor(), T.Normalize(MEAN, STD)])1314def block(cin, cout):15 return nn.Sequential(16 nn.Conv2d(cin, cout, 3, padding=1, bias=False), nn.BatchNorm2d(cout), nn.ReLU(),17 nn.Conv2d(cout, cout, 3, padding=1, bias=False), nn.BatchNorm2d(cout), nn.ReLU(),18 nn.MaxPool2d(2),19 )2021class SmallVGG(nn.Module):22 def __init__(self, n_classes=10):23 super().__init__()24 self.features = nn.Sequential(block(3, 64), block(64, 128), block(128, 256))25 self.head = nn.Sequential(nn.Flatten(), nn.Dropout(0.3), nn.Linear(256 * 4 * 4, n_classes))26 def forward(self, x):27 return self.head(self.features(x))2829model = SmallVGG()30opt = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9, weight_decay=5e-4, nesterov=True)31sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=0.1, total_steps=30 * len(train_loader))32criterion = nn.CrossEntropyLoss(label_smoothing=0.1)Two things to get right, because they are the ones that bite later rather than now.
Record the normalisation constants somewhere the serving code reads. Put MEAN and STD into metadata.json, and have the server load them from there rather than hard-coding them a second time. A mismatch between training and serving normalisation is the single most common cause of "the model scored 89% offline and 61% in production", and it produces no error at all.
Save a per-class breakdown, not just overall accuracy. CIFAR-10's cat and dog classes routinely sit ten points below the rest. If you only record 88.3% overall, you will not notice when a later change drops cats to 61%.
1from sklearn.metrics import classification_report, confusion_matrix2report = classification_report(y_true, y_pred, target_names=CLASSES, output_dict=True)3json.dump({"overall_acc": acc, "per_class": report,4 "confusion": confusion_matrix(y_true, y_pred).tolist()},5 open("models/metrics.json", "w"), indent=2)Phase 2 — Export and verify the artefact
The trained nn.Module is not deployable: loading it requires the SmallVGG class definition. Export to a self-contained artefact and prove it computes the same thing.
1# training/export.py2import torch, numpy as np, onnx, onnxruntime as ort, hashlib, json, subprocess34model.eval() # critical: freezes dropout and batch-norm5dummy = torch.randn(2, 3, 32, 32) # batch 2: size-1 dims cannot be made dynamic67traced = torch.jit.trace(model, dummy) # TorchScript: deprecated, still works8traced.save("models/model.pt")910torch.onnx.export(model, (dummy,), "models/model.onnx",11 input_names=["pixels"], output_names=["logits"],12 dynamic_shapes={"x": {0: torch.export.Dim("batch")}},13 opset_version=18)1415onnx.checker.check_model(onnx.load("models/model.onnx"))16sess = ort.InferenceSession("models/model.onnx", providers=["CPUExecutionProvider"])1718worst = 0.019for batch in (1, 4, 16): # exercise the dynamic axis20 x = np.random.randn(batch, 3, 32, 32).astype(np.float32)21 with torch.no_grad():22 expected = model(torch.from_numpy(x)).numpy()23 actual = sess.run(None, {"pixels": x})[0]24 worst = max(worst, float(np.abs(expected - actual).max()))2526assert worst < 1e-4, f"ONNX diverges from PyTorch by {worst:.3e}"2728digest = hashlib.sha256(open("models/model.pt", "rb").read()).hexdigest()29json.dump({30 "version": "1.0.0",31 "sha256": digest,32 "opset_version": 18,33 "exported_from": torch.__version__,34 "max_abs_diff_vs_pytorch": worst,35 "normalisation": {"mean": list(MEAN), "std": list(STD)},36 "classes": CLASSES,37 "input_shape": ["batch", 3, 32, 32],38 "git_commit": subprocess.check_output(["git", "rev-parse", "HEAD"]).decode().strip(),39}, open("models/metadata.json", "w"), indent=2)Three parts of that file earn their keep later. model.eval() before tracing: export in training mode and dropout stays active in production, randomly zeroing 30% of your features on every request. The batch-size loop: exporting with one batch size and never testing another is how a dynamic axis bug reaches production. And metadata.json: when a prediction looks wrong in three months, the digest tells you which bytes produced it and the normalisation block tells you what preprocessing they expect.
Then freeze a golden set — twelve real images and the outputs the verified model produces for them:
1inputs = np.stack([test_tfm(load(p)).numpy() for p in GOLDEN_PATHS])2np.savez("tests/golden.npz", inputs=inputs,3 logits=sess.run(None, {"pixels": inputs})[0],4 labels=np.array(GOLDEN_LABELS))This one file catches more later regressions than any other test you will write: label reordering, changed preprocessing, a corrupted artefact, or the wrong model baked into an image.
Phase 3 — Serve it
1# app/main.py2from fastapi import FastAPI, HTTPException, UploadFile, File, Depends3from contextlib import asynccontextmanager4import torch, io, time, json, uuid5from PIL import Image6import torchvision.transforms as T78STATE = {}9MAX_BYTES = 5 * 1024 * 10241011@asynccontextmanager12async def lifespan(app: FastAPI):13 meta = json.load(open("models/metadata.json"))14 torch.set_num_threads(1)15 m = torch.jit.load("models/model.pt", map_location="cpu"); m.eval()16 STATE.update(17 model=m, classes=meta["classes"], version=meta["version"],18 tfm=T.Compose([T.Resize((32, 32)), T.ToTensor(),19 T.Normalize(meta["normalisation"]["mean"],20 meta["normalisation"]["std"])]),21 warmup=torch.zeros(1, 3, 32, 32),22 )23 with torch.no_grad():24 m(STATE["warmup"]) # first call is slow; pay it before traffic25 yield26 STATE.clear()2728app = FastAPI(title="CIFAR-10 Classifier", version="1.0.0", lifespan=lifespan)2930@app.get("/healthz")31def healthz():32 return {"status": "alive"}3334@app.get("/readyz")35def readyz():36 try:37 with torch.no_grad():38 STATE["model"](STATE["warmup"])39 return {"ready": True, "model_version": STATE["version"]}40 except Exception as e:41 raise HTTPException(503, f"not ready: {e}")4243@app.post("/predict", response_model=Prediction)44def predict(file: UploadFile = File(...), _=Depends(check_rate)):45 if file.content_type not in {"image/jpeg", "image/png"}:46 raise HTTPException(415, f"unsupported content type: {file.content_type}")47 raw = file.file.read(MAX_BYTES + 1)48 if len(raw) > MAX_BYTES:49 raise HTTPException(413, "image exceeds 5 MB")50 try:51 img = Image.open(io.BytesIO(raw)).convert("RGB")52 except Exception:53 raise HTTPException(422, "file is not a decodable image")5455 t0 = time.perf_counter()56 with torch.no_grad():57 probs = torch.softmax(STATE["model"](STATE["tfm"](img).unsqueeze(0)), dim=1)[0]58 ms = (time.perf_counter() - t0) * 100059 idx = int(probs.argmax())6061 top3 = torch.topk(probs, 3)62 return Prediction(63 label=STATE["classes"][idx],64 confidence=round(float(probs[idx]), 4),65 top3=[{"label": STATE["classes"][int(i)], "p": round(float(p), 4)}66 for p, i in zip(top3.values, top3.indices)],67 model_version=STATE["version"],68 latency_ms=round(ms, 2),69 )Four decisions in that file are worth defending if someone reviews your work.
Why TorchScript here at all. It is deprecated in current PyTorch but still works, and it keeps this server in the same PyTorch code you trained with. If you are starting fresh, serving model.onnx with ONNX Runtime is the better long-term choice: the handler changes by a few lines (an InferenceSession instead of torch.jit.load, NumPy instead of tensors) and PyTorch drops out of the image entirely.
The model loads in lifespan, not in the handler. Loading inside predict makes every request pay deserialisation cost, which turns a 40 ms endpoint into a 400 ms one and looks like a mysterious performance problem.
def, not async def. The body contains a blocking PyTorch call and no await. Declared async, that call blocks the event loop for the whole forward pass, so ten simultaneous requests are served strictly one after another and your health check queues behind them.
The warm-up call. The first inference after load includes allocation and kernel selection and can be five times slower than steady state. Doing it before yield means the first real user does not absorb it — and it also makes /readyz meaningful, since a model that cannot run a forward pass is not ready.
Reading MAX_BYTES + 1. Checking size after reading the whole file means a 2 GB upload is already in memory when you reject it.
Phase 4 — Test it at three levels
| Level | What it covers | Model | Runtime |
|---|---|---|---|
| Unit | Preprocessing shapes, schema validation, label mapping | Mocked | Seconds |
| Golden | The artefact still produces the recorded outputs | Real, in-process | Seconds |
| Integration | The built container answers real HTTP correctly | Real, in a container | Minutes |
1# tests/integration/test_api.py2import requests, pathlib, pytest, numpy as np34def test_ready(base_url):5 assert requests.get(f"{base_url}/readyz", timeout=10).json()["ready"] is True67@pytest.mark.parametrize("path,expected", [8 ("tests/fixtures/cat_01.jpg", "cat"),9 ("tests/fixtures/ship_03.jpg", "ship"),10 ("tests/fixtures/truck_07.jpg", "truck"),11])12def test_known_images(base_url, path, expected):13 with open(path, "rb") as f:14 r = requests.post(f"{base_url}/predict", files={"file": f}, timeout=30)15 assert r.status_code == 20016 body = r.json()17 assert body["label"] == expected18 assert body["model_version"] == "1.0.0" # right artefact is inside the image1920def test_rejects_text_file(base_url):21 r = requests.post(f"{base_url}/predict",22 files={"file": ("a.txt", b"hello", "text/plain")}, timeout=10)23 assert r.status_code == 415 # not 5002425def test_rejects_corrupt_image(base_url):26 r = requests.post(f"{base_url}/predict",27 files={"file": ("a.jpg", b"\xff\xd8notanimage", "image/jpeg")},28 timeout=10)29 assert r.status_code == 4223031def test_p95_latency(base_url):32 img = pathlib.Path("tests/fixtures/cat_01.jpg").read_bytes()33 for _ in range(10): # warm up34 requests.post(f"{base_url}/predict", files={"file": ("a.jpg", img)})35 times = []36 for _ in range(200):37 r = requests.post(f"{base_url}/predict", files={"file": ("a.jpg", img)})38 times.append(r.json()["latency_ms"])39 p95 = float(np.percentile(times, 95))40 assert p95 < 100, f"p95 {p95:.1f} ms exceeds budget"The test_rejects_text_file and test_rejects_corrupt_image cases are the ones people skip and the ones that matter most for a service someone else operates. A 500 means "we are broken"; a 415 or 422 means "your input was wrong". If every failure is a 500, the operator cannot distinguish their mistake from yours, and your error dashboard is noise.
Phase 5 — Containerise and automate
1FROM python:3.11-slim-bookworm23RUN apt-get update \4 && apt-get install -y --no-install-recommends libgomp1 libjpeg62-turbo curl \5 && rm -rf /var/lib/apt/lists/*67ENV PYTHONUNBUFFERED=1 PYTHONDONTWRITEBYTECODE=1 OMP_NUM_THREADS=18WORKDIR /app910COPY requirements.txt .11RUN pip install --no-cache-dir -r requirements.txt1213COPY models/ /app/models/14COPY app/ /app/app/1516RUN useradd --create-home --uid 10001 appuser && chown -R appuser:appuser /app17USER appuser1819EXPOSE 800020HEALTHCHECK --interval=30s --timeout=3s --start-period=40s --retries=3 \21 CMD curl -fsS http://localhost:8000/readyz || exit 12223CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "2"]# requirements.txt — the CPU index is what keeps the image under 1.5 GB--extra-index-url https://download.pytorch.org/whl/cputorch==2.14.0+cputorchvision==0.29.0+cpufastapi==0.141.1uvicorn[standard]==0.53.0pillow==12.3.0python-multipart==0.0.32The +cpu wheels are the difference between roughly 5.5 GB and roughly 1.1 GB. The default Linux install of PyTorch pulls in about 4.5 GB of CUDA libraries that a CPU-only container will never execute. If your image is over budget, check this before anything else.
The dependency ordering — requirements, then models, then app/ — is what makes rebuilds fast. Docker rebuilds every layer after the first changed one, so putting application code last means a code-only change rebuilds in seconds instead of reinstalling PyTorch.
1name: ci2on: [push, pull_request]3jobs:4 build-and-test:5 runs-on: ubuntu-latest6 steps:7 - uses: actions/checkout@v58 with: { lfs: true }9 - uses: actions/setup-python@v610 with: { python-version: "3.11", cache: pip }11 - run: pip install -r requirements.txt -r requirements-dev.txt12 - run: pytest tests/unit -q13 - run: python scripts/check_golden.py # artefact still matches14 - run: docker build -t cifar-api:test .15 - name: Image size budget16 run: |17 BYTES=$(docker image inspect cifar-api:test --format '{{.Size}}')18 echo "image size: $((BYTES / 1000000)) MB"19 test "$BYTES" -lt 150000000020 - name: Start container and wait for readiness21 run: |22 docker run -d --name api -p 8000:8000 --memory 2g --cpus 2 cifar-api:test23 for i in $(seq 1 30); do24 curl -fsS http://localhost:8000/readyz && break25 sleep 226 done27 - run: pytest tests/integration -q --base-url http://localhost:800028 - if: failure()29 run: docker logs apiThe --memory 2g --cpus 2 flags matter: testing on an unconstrained GitHub runner tells you nothing about behaviour on the two-core box this will actually run on. The readiness loop rather than a fixed sleep is what keeps the pipeline from being flaky.
Where to deploy it
| Option | Effort | Cost | Fits |
|---|---|---|---|
| Docker Compose on the customer's own machine | An hour | None | The conveyor scenario: local network, no cloud dependency |
| Hugging Face Spaces (Gradio) | An hour | Free tier | Showing it to stakeholders; sleeps when idle |
| A small cloud VM behind a proxy | Half a day | Low, fixed | Steady low traffic, full control |
| Serverless container (Cloud Run, Lambda) | Half a day | Near zero when idle | Spiky or very low traffic; tolerate 4–10 s cold starts |
| Managed endpoint (SageMaker, Vertex AI) | A day | Instance-hours, always on | Steady traffic above a couple of requests per second |
For the scenario at the top of this lesson, Compose on a machine at the factory is the right answer, and reaching for a managed cloud endpoint would be a mistake — it adds a network dependency the line cannot tolerate and a monthly bill for two requests per second. Deploy at least two of these anyway, because the second one is where you discover which of your assumptions were about your laptop rather than about your software.
What goes wrong, and what it looks like
| Symptom | Likely cause | Fix |
|---|---|---|
| Accuracy far worse in the API than in the notebook | Serving normalisation differs from training | Load mean and std from metadata.json; assert on the golden set |
| Predictions differ between two identical requests | model.eval() missing before export; dropout still active | Re-export in eval mode; verify determinism in a test |
| Container restarts in a loop | Health check fires before the model finishes loading | Set --start-period=40s; use /healthz for liveness |
| Every request takes several hundred ms | Model loaded inside the handler | Move it to lifespan; log once at load and count the lines |
| p99 far worse than p50 under concurrency | async def around a blocking call, or thread oversubscription | Use plain def; set OMP_NUM_THREADS=1 with multiple workers |
| Image is 5 GB or more | Default PyTorch wheel with CUDA | Install from the CPU index; check with docker history |
| Memory climbs steadily under load | Missing torch.no_grad(), retaining activations | Wrap inference; watch RSS over a ten-minute load test |
| Model file loads as 130 bytes in CI | git-lfs not fetched in the workflow | Add with: { lfs: true } to the checkout step |
| Works locally, 404s in the container | Relative model path resolved against a different working directory | Use absolute paths or set WORKDIR explicitly |
Almost every failure in this list produces a working service that is quietly wrong, rather than an error — which is why the golden set and the integration tests are not optional extras.
How to know you have actually finished
The honest test is not whether the endpoint returns 200. It is whether someone else can run it, and whether you can diagnose it when it misbehaves. Three concrete exercises settle that.
Hand it over. Give a colleague only the repository URL and the README. They should get a working service with one command and no questions. Every question they ask is a gap in your README, and the fix is to write the answer down rather than to explain it.
Break it deliberately. Delete models/model.pt and start the container — it should fail loudly at startup, not serve 500s to users. Send a 50 MB file — it should return 413 quickly, not consume all your memory. Kill the container mid-request — the client should get a clean connection error, not a hang. Point it at the wrong metadata file — /readyz should refuse to report ready.
Reconstruct a past request. Run a load test, then pick one request from an hour ago and answer: which model version served it, how long inference took, what it predicted, and with what confidence. If you cannot answer all four from the logs in under two minutes, your logging is not finished, and you will find that out during an incident instead.
Anything beyond that — quantizing the model to shrink it further, adding a Redis cache, wiring up Prometheus, running a canary deployment — is a genuine improvement, but each one is optional. The three exercises above are not. A service that passes them is deployable by someone who is not you, which is the only definition of "deployed" that survives contact with a real organisation.