Course Content
AI System Design and Architecture
3 sections · 7 lessons
Distributed Model Serving — Batching and Sharding
An engineer benchmarks a BERT-base classifier on a T4 GPU. One request in, one out: 7.2 ms. She multiplies: 1000/7.2 = 139 requests per second per GPU. Traffic forecast is 600 req/s, so she provisions five GPUs, adds one for headroom, and ships six.
The service works. It is also costing about four times what it should, and its p99 latency is worse than a single-GPU deployment would have delivered. The reason is that a GPU does not process one request in 7.2 ms — it processes one batch in 7.2 ms, and a batch of one wastes nearly all of the silicon. Feed the same card 13 requests at a time and it handles 600 req/s on its own.
Distributed model serving is the discipline of getting inference off a laptop and into a system that survives real traffic. Almost every important decision in it comes down to arithmetic you can do on paper before you provision anything.
What "serving" actually involves
Serving a model is not "calling predict() over HTTP". A production serving layer owns six responsibilities, and skipping any one of them shows up as an incident:
- Loading and holding weights — expensive, so it happens once at startup, not per request.
- Preprocessing — turning JSON or bytes into exactly the tensor shape the model was trained on. Version skew here silently destroys accuracy.
- Scheduling — deciding which requests execute together and in what order. This is where most of the performance lives.
- Execution — the forward pass on CPU, GPU, or accelerator.
- Postprocessing — logits to labels, tokens to text, plus thresholds and business rules.
- Observability — per-stage timings, input distributions, output distributions, model version on every response.
It becomes distributed the moment one machine cannot do it: because throughput exceeds one device, because the weights exceed one device's memory, or because you need the service to survive a machine dying.
The batching arithmetic that decides your bill
Model a GPU forward pass as a fixed cost plus a marginal cost per item. For BERT-base at 128 tokens on a T4, measured:
The 6 ms is kernel launches, memory transfers and framework overhead — paid once regardless of batch size. The 1.2 ms per sequence is the actual arithmetic. Now compute throughput b/T(b):
| Batch size b | T(b) | Throughput | Cost per request (GPU-ms) |
|---|---|---|---|
| 1 | 7.2 ms | 139 req/s | 7.20 |
| 4 | 10.8 ms | 370 req/s | 2.70 |
| 8 | 15.6 ms | 513 req/s | 1.95 |
| 16 | 25.2 ms | 635 req/s | 1.58 |
| 32 | 44.4 ms | 721 req/s | 1.39 |
| 64 | 82.8 ms | 773 req/s | 1.29 |
| 128 | 159.6 ms | 802 req/s | 1.25 |
Batching 1 to 16 multiplies throughput by 4.6×. Batching 16 to 128 adds only another 26%, while multiplying single-batch execution time by 6.3×. The returns collapse exactly where the fixed cost stops dominating, and that is the batch size you want.
Dynamic batching finds the right size by itself
You do not have to guess. A dynamic batching server collects whatever arrived while the GPU was busy and runs that as the next batch. In steady state at arrival rate λ, the batch size is the number of arrivals during one execution:
With λ = 600 req/s = 0.6 req/ms: b=3.6+0.72b, so 0.28b=3.6 and b=12.9. Execution time is T=6+1.2(12.9)=21.4 ms, and throughput is 12.9/21.4=0.60 req/ms = 600 req/s. It balances itself.
Average latency for a request is the time it waits for the current batch to finish collecting — uniformly distributed over the 21.4 ms window, so 10.7 ms on average — plus its own batch's 21.4 ms execution: 32.1 ms on one GPU.
Compare that with the six-GPU no-batching deployment. Each of five active replicas takes 120 req/s with μ = 139 req/s, so ρ = 0.864 and queue wait is ρ/(μ−λ)=0.864/18.9=45.7 ms, for a total of 52.9 ms.
| No batching | Dynamic batching | |
|---|---|---|
| GPUs needed for 600 req/s | 5 (6 with headroom) | 1 (2 with headroom) |
| Average latency | 52.9 ms | 32.1 ms |
| Hourly cost at USD 0.526/GPU | USD 3.16 | USD 1.05 |
| Cost per million requests | USD 1.46 | USD 0.49 |
Dynamic batching is the single highest-leverage change in model serving: it usually cuts cost by 3–5× and improves latency at the same time. Almost nothing else does both.
The one control you must set is the maximum queue delay — how long the scheduler will wait for a batch to fill when traffic is light. Set it near your latency budget's slack, typically 5–20 ms. Set it too high and quiet-period requests sit waiting for company that never arrives.
Three serving patterns
| Batch prediction | Real-time | Hybrid (precompute + fallback) | |
|---|---|---|---|
| Trigger | Schedule (nightly, hourly) | User request | Schedule, with on-demand for misses |
| Latency seen by user | Zero — it is a lookup | Full inference cost | Lookup on hit, inference on miss |
| Freshness | Stale by up to one cycle | Always current | Fresh where it matters |
| GPU efficiency | Highest — huge batches, spot instances | Lowest — must be provisioned for peak | High |
| Fails when | Input is unknown until request time | Traffic spikes faster than you can scale | Hit rate is low |
| Typical use | Recommendations for known users, risk scores, nightly embeddings | Fraud checks, chat, search ranking, moderation | Product catalogues, top-N recommendations |
The hybrid pattern deserves more credit than it gets. Precompute predictions for the 5% of entities that receive 80% of traffic, serve those from a key-value store in 2 ms, and fall back to live inference for the tail. If 70% of requests hit precomputed results, effective latency at a 45 ms model is 0.7×2+0.3×45=1.4+13.5=14.9 ms, and your live GPU fleet shrinks by 70%.
How to package it
Container-based serving
Your own web framework, your own model loading, your own batching. Total control, and total responsibility.
1from contextlib import asynccontextmanager2from fastapi import FastAPI3import torch4from torch.export.passes import move_to_device_pass56STATE = {}78@asynccontextmanager9async def lifespan(app: FastAPI):10 ep = torch.export.load("/models/classifier_v7.pt2") # from torch.export.save11 m = move_to_device_pass(ep, "cuda").module()12 with torch.inference_mode(): # warm the CUDA kernels13 m(torch.zeros(8, 128, dtype=torch.long, device="cuda"))14 STATE["model"], STATE["version"] = m, "v7"15 yield16 STATE.clear()1718app = FastAPI(lifespan=lifespan)1920@app.get("/healthz") # liveness: is the process alive?21def healthz(): return {"ok": True}2223@app.get("/readyz") # readiness: are the weights actually loaded?24def readyz():25 return ({"ready": True, "model": STATE["version"]} if "model" in STATE26 else ({"ready": False}, 503))The model file is a torch.export artefact, exported in eval mode with a dynamic batch dimension so the server can run any batch size. (TorchScript's torch.jit is deprecated in current PyTorch 2.x releases; torch.export is its replacement.)
The separation of /healthz from /readyz matters more than it looks. If the load balancer's health check hits an endpoint that returns 200 before the weights are resident, it sends traffic to a replica that will time out for the next 45 seconds. Every scale-up event then produces a burst of errors, and the postmortem blames the autoscaler.
Specialised serving frameworks
| Framework | What it is good at | Cost of adopting it |
|---|---|---|
| TensorFlow Serving | Versioned SavedModel loading, gRPC, hot model swap without restart | TensorFlow only; awkward custom pre/postprocessing |
| NVIDIA Triton (now branded Dynamo-Triton) | Multiple frameworks on one server, dynamic batching, concurrent model instances per GPU, ensemble graphs | Configuration is verbose; a real learning curve |
| TorchServe | Native PyTorch handlers, straightforward custom logic | No longer maintained: the repository was archived in August 2025, so avoid it for new work |
| BentoML | Python-first packaging, "bentos" bundling code plus weights plus environment, good developer ergonomics | Another abstraction between you and the runtime |
| vLLM / SGLang | LLM-specific: paged KV cache, continuous batching, huge throughput gains for generation | Generative transformers only. (Hugging Face TGI is now in maintenance mode and points users to these.) |
The feature that justifies Triton for mixed workloads is concurrent model instances: two or four copies of a small model on one GPU, so that while instance A is in a memory-bound phase, instance B occupies the compute units. On a model that leaves the GPU 40% idle during data movement, running three instances can lift utilisation from 60% to over 90% without touching the model.
Serverless, and exactly when it is cheaper
Serverless removes idle cost but adds cold starts and a much higher per-request price. Do the crossover calculation rather than arguing about it.
A 3 GB function running 500 ms per request costs 1.5 GB-seconds. At AWS Lambda's list price of USD 0.0000166667 per GB-second (as of September 2026, ignoring the small per-request fee) that is USD 0.000025 per request. An always-on pair of c6i.2xlarge instances costs 2 × 0.34 = USD 0.68/hr = USD 496 per month and comfortably serves this workload. Break-even volume is
Below about 7 req/s, serverless wins. Above it, always-on wins and the gap widens fast: at 100 req/s serverless costs roughly USD 6,570 a month against USD 496, a 13× penalty. Add that a 3 GB model pulled from network storage takes 6–10 seconds to cold start, and at low traffic — precisely where serverless is cheapest — most requests pay it.
Serverless is neither cheap nor expensive — it has a crossover point, and you can compute it in one line before the argument starts.
Scaling in four directions
Horizontal: more replicas
The default and usually the right one. Each replica holds a full copy of the weights and serves independent requests. Capacity is linear in replica count; the only constraints are memory per replica and the load balancer's ability to spread work evenly.
┌───────────────┐ requests ───────►│ LB: least │ │ outstanding │ └───┬───┬───┬───┘ ┌──────────┘ │ └──────────┐ ▼ ▼ ▼ ┌────────────┐ ┌────────────┐ ┌────────────┐ │ replica 1 │ │ replica 2 │ │ replica 3 │ │ full model │ │ full model │ │ full model │ │ batch ≤ 32 │ │ batch ≤ 32 │ │ batch ≤ 32 │ └────────────┘ └────────────┘ └────────────┘Vertical: a bigger machine
Useful in exactly one situation: the model does not fit, or a single request is too slow on current hardware. It is not a throughput strategy — a card with 2× the memory rarely gives 2× the throughput, and it gives you no additional failure domain. Two small instances beat one large one for availability every time.
Sharding: one model across several GPUs
Sometimes the weights simply do not fit. Llama-3-70B in fp16 is 70×109×2=140 GB of parameters. The most widely deployed data-centre GPU, the 80 GB H100, cannot hold it. Newer cards with 141–192 GB can hold the weights but leave little room for the KV cache described below, so on most fleets the model must be split across devices.
Tensor parallelism splits each layer's matrices across devices, so all GPUs work on every token and synchronise with a collective operation twice per layer. It keeps latency low and needs a fast interconnect — NVLink inside one machine, not Ethernet between machines.
Pipeline parallelism gives each device a contiguous block of layers and passes activations along. It tolerates slower links but introduces bubbles: with 4 stages, a single request leaves 3 of 4 devices idle at any moment, so you must keep several micro-batches in flight to fill the pipe.
Tensor parallel (TP=4) Pipeline parallel (PP=4) ┌────┬────┬────┬────┐ ┌────┐ ┌────┐ ┌────┐ ┌────┐ │GPU0│GPU1│GPU2│GPU3│ │L0-19│►│L20-39│►│L40-59│►│L60-79│ │ 1/4 of every layer │ └────┘ └────┘ └────┘ └────┘ └────┴────┴────┴────┘ latency = sum of stages all-reduce ×2 per layer bubbles unless micro-batchedMemory planning matters as much as the split. Beyond weights, generation needs a KV cache. For a model with 80 layers, 8 key/value heads and head dimension 128 in fp16, each token costs 2×80×8×128×2=327,680 bytes = 320 KiB. A 4,096-token conversation therefore holds 1.25 GiB of cache. On four 80 GB cards you have 320 GB total, 140 GB goes to weights, leaving about 167 GiB — roughly 130 concurrent 4k conversations, before activations. That number, not FLOPs, is what caps concurrency on an LLM endpoint.
Ensembles: several models per request
Running three models and combining outputs triples cost unless you are careful. The cheap version is cascading: run a small fast model first, and escalate to the expensive one only when confidence is low. If the small model handles 85% of traffic at 3 ms and the large one takes 60 ms, average latency is 0.85×3+0.15×63=2.55+9.45=12 ms and GPU cost falls by roughly 80% versus always running the large model.
Caching in front of the model
Inference is deterministic for a fixed model version and fixed input, which makes it unusually cacheable.
| Cache type | Key | Typical hit rate | Watch out for |
|---|---|---|---|
| Query result | hash(normalised input + model version) | 15–45%, higher with skewed traffic | Never omit model version from the key |
| Embedding | hash(document id + encoder version) | 80–99% for stable corpora | Storage grows with corpus size |
| Partial / prefix | hash(shared prompt prefix) | Very high for templated prompts | Needs runtime support (KV prefix reuse) |
| Negative | hash(input) for known-bad inputs | Small but valuable | Short TTL, or you cache a transient failure |
Quantify a modest 40% hit rate on a 45 ms model with a 2 ms cache: effective latency is 0.4×2+0.6×45=27.8 ms, a 38% improvement, and GPU load drops 40%, so a 10-replica fleet becomes 6. On the earlier batched deployment at USD 0.49 per million requests, that is USD 0.29 per million.
The failure mode is stale results after a model update. The fix is not invalidation — it is putting the model version in the key. Deploy v8 and every v7 key becomes unreachable and ages out naturally; roll back to v7 and its cache is still warm. Compare that with a flush-on-deploy strategy, where every deployment sends 100% of traffic to cold GPUs at once, which is how a routine release becomes an outage.
Knowing whether it works
Serving metrics fall into two families, and teams reliably instrument only the first.
| System health | Model health |
|---|---|
| Latency p50 / p95 / p99, split by stage | Distribution of predicted classes over time |
| Requests per second and batch size histogram | Mean confidence and its drift |
| Error rate by status code | Input feature distributions vs training data |
| GPU utilisation and memory | Rate of fallback / low-confidence escalation |
| Cache hit rate | Delayed accuracy against labels when they arrive |
A model whose average confidence drifts from 0.91 to 0.68 over three weeks is failing, and every system metric will stay green throughout. Emit the model version and the confidence on every response, log a sampled 1% of inputs with their outputs, and compare this week's distribution with training. That comparison is the only early warning you will get.
Split latency by stage, always. "p99 is 210 ms" is not actionable; "p99 is 210 ms, of which 12 ms is preprocessing, 22 ms is batch wait, 160 ms is inference and 16 ms is postprocessing" tells you precisely where to spend the next day.
Putting the numbers to work
When you size a serving deployment, do it in this order and you will rarely be wrong. First measure T(b) for at least four batch sizes on the hardware you intend to buy — that one table decides your replica count, your cost, and your latency at once. Second, pick the batch size where marginal throughput gain drops below about 10%, and set the maximum queue delay to whatever latency slack you have. Third, divide peak traffic by the throughput at that batch size, then divide again by 0.6 so you are provisioned at 60% utilisation, because the 1/(1−ρ) shape of queue delay makes 90% utilisation feel broken. Fourth, check whether the weights fit; if they do not, you are in sharding territory and the KV cache calculation caps your concurrency.
Two failure modes are worth naming because they are so common. Benchmarking at batch 1 and provisioning for it is the opening scenario, and it costs a multiple of the correct bill. Health checks that pass before weights are loaded turn every scale-up into an error spike and make autoscaling look dangerous when it is not — separate liveness from readiness and gate readiness on a completed warm-up inference.
Get those right and the rest is ordinary engineering. Get them wrong and no amount of infrastructure sophistication will save the deployment, because you will be scaling a system whose fundamental unit of work you have mismeasured.