AI System Design and Architecture

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 batching does to one T4 GPU7.2 ms139 req/s7.2025.2 ms635 req/s1.5844.4 ms721 req/s1.39159.6 ms802 req/s1.25GPU timeThroughputGPU-ms per requestBatch of 1Batch of 16Batch of 32Batch of 128Six GPUs sized from the batch-of-1 number serve 600 req/s; one GPU batching about 13 at a time does the same.
A GPU is idle most of a single request, so batching spends latency you can afford to buy throughput you were renting.

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:

T(b)=6+1.2b msT(b) = 6 + 1.2b \text{ ms}

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)b/T(b):

Batch size bT(b)ThroughputCost per request (GPU-ms)
17.2 ms139 req/s7.20
410.8 ms370 req/s2.70
815.6 ms513 req/s1.95
1625.2 ms635 req/s1.58
3244.4 ms721 req/s1.39
6482.8 ms773 req/s1.29
128159.6 ms802 req/s1.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:

b=λ⋅T(b)=λ(6+1.2b)b = \lambda \cdot T(b) = \lambda(6 + 1.2b)

With λ = 600 req/s = 0.6 req/ms: b=3.6+0.72bb = 3.6 + 0.72b, so 0.28b=3.60.28b = 3.6 and b=12.9b = 12.9. Execution time is T=6+1.2(12.9)=21.4T = 6 + 1.2(12.9) = 21.4 ms, and throughput is 12.9/21.4=0.6012.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\rho/(\mu - \lambda) = 0.864/18.9 = 45.7 ms, for a total of 52.9 ms.

No batchingDynamic batching
GPUs needed for 600 req/s5 (6 with headroom)1 (2 with headroom)
Average latency52.9 ms32.1 ms
Hourly cost at USD 0.526/GPUUSD 3.16USD 1.05
Cost per million requestsUSD 1.46USD 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 predictionReal-timeHybrid (precompute + fallback)
TriggerSchedule (nightly, hourly)User requestSchedule, with on-demand for misses
Latency seen by userZero — it is a lookupFull inference costLookup on hit, inference on miss
FreshnessStale by up to one cycleAlways currentFresh where it matters
GPU efficiencyHighest — huge batches, spot instancesLowest — must be provisioned for peakHigh
Fails whenInput is unknown until request timeTraffic spikes faster than you can scaleHit rate is low
Typical useRecommendations for known users, risk scores, nightly embeddingsFraud checks, chat, search ranking, moderationProduct 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.90.7 \times 2 + 0.3 \times 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.

Python
from contextlib import asynccontextmanagerfrom fastapi import FastAPIimport torchfrom torch.export.passes import move_to_device_passSTATE = {}@asynccontextmanagerasync def lifespan(app: FastAPI):    ep = torch.export.load("/models/classifier_v7.pt2")   # from torch.export.save    m = move_to_device_pass(ep, "cuda").module()    with torch.inference_mode():          # warm the CUDA kernels        m(torch.zeros(8, 128, dtype=torch.long, device="cuda"))    STATE["model"], STATE["version"] = m, "v7"    yield    STATE.clear()app = FastAPI(lifespan=lifespan)@app.get("/healthz")      # liveness: is the process alive?def healthz(): return {"ok": True}@app.get("/readyz")       # readiness: are the weights actually loaded?def readyz():    return ({"ready": True, "model": STATE["version"]} if "model" in STATE            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

FrameworkWhat it is good atCost of adopting it
TensorFlow ServingVersioned SavedModel loading, gRPC, hot model swap without restartTensorFlow only; awkward custom pre/postprocessing
NVIDIA Triton (now branded Dynamo-Triton)Multiple frameworks on one server, dynamic batching, concurrent model instances per GPU, ensemble graphsConfiguration is verbose; a real learning curve
TorchServeNative PyTorch handlers, straightforward custom logicNo longer maintained: the repository was archived in August 2025, so avoid it for new work
BentoMLPython-first packaging, "bentos" bundling code plus weights plus environment, good developer ergonomicsAnother abstraction between you and the runtime
vLLM / SGLangLLM-specific: paged KV cache, continuous batching, huge throughput gains for generationGenerative 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

4960.000025=19,840,000 requests/month=7.6 req/s\frac{496}{0.000025} = 19{,}840{,}000 \text{ requests/month} = 7.6 \text{ req/s}

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.

Text
                    ┌───────────────┐   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=14070 \times 10^9 \times 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.

Text
  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-batched

Memory 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,6802 \times 80 \times 8 \times 128 \times 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=120.85 \times 3 + 0.15 \times 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 typeKeyTypical hit rateWatch out for
Query resulthash(normalised input + model version)15–45%, higher with skewed trafficNever omit model version from the key
Embeddinghash(document id + encoder version)80–99% for stable corporaStorage grows with corpus size
Partial / prefixhash(shared prompt prefix)Very high for templated promptsNeeds runtime support (KV prefix reuse)
Negativehash(input) for known-bad inputsSmall but valuableShort 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.80.4 \times 2 + 0.6 \times 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 healthModel health
Latency p50 / p95 / p99, split by stageDistribution of predicted classes over time
Requests per second and batch size histogramMean confidence and its drift
Error rate by status codeInput feature distributions vs training data
GPU utilisation and memoryRate of fallback / low-confidence escalation
Cache hit rateDelayed 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)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.