Serving an image classifier returns a handful of numbers per image. Serving a segmentation model returns a label for every pixel, and sometimes many overlapping instance masks on top of that. The output can be as large as the input, post-processing can cost more than the model, and a one-pixel mistake in the resize bookkeeping shifts every mask in production while every offline metric still looks fine.
This article covers serving three families of segmentation models: semantic segmentation (one class per pixel, for example SegFormer or DeepLab), instance and panoptic segmentation (a set of masks with classes and scores, for example Mask R-CNN or Mask2Former), and promptable segmentation (the Segment Anything family, where a click, box or text phrase picks the object). It assumes you already know the basics of GPU image decode, TensorRT engines and Triton ensembles from Image Classification Serving, and spends its words on what is different when the answer is a mask: output arithmetic, GPU post-processing, mask encoding, the encoder and decoder split, tiling large images, and the failure modes that only segmentation has.
Why segmentation serving is different
Three properties drive every design decision. First, the output tensor scales with image area and class count. A classifier head emits 1,000 logits; a semantic head for 150 classes at quarter resolution on a 512 by 512 input emits 128 x 128 x 150 = 2,457,600 logits. Second, the answer must be mapped back to the original pixel grid, which means you have to remember exactly how you resized, padded and cropped the input. Third, the client rarely wants raw logits. It wants a compact mask format, polygons for an annotation tool, or per-instance boxes and areas, so encoding is part of the serving path rather than a client concern.
So a segmentation server is a four-stage pipeline: preprocessing (GPU, cheap), the network (GPU, heavy), post-processing (GPU if designed well, slow CPU if not) and encoding (CPU, often the real bottleneck). Profile all four before tuning one.
Output arithmetic decides the design
Do the output arithmetic before you choose where post-processing runs. Take a semantic model with 150 classes (the ADE20K label set) that predicts at one quarter of the input resolution, served at 512 by 512 with a batch of 16.
| Tensor | Shape per image | Bytes per image | Batch of 16 |
|---|---|---|---|
| Raw logits, fp32 | 150 x 128 x 128 | 9.8 MB | 157 MB |
| Upsampled logits, fp32 | 150 x 512 x 512 | 157 MB | 2.5 GB |
| Upsampled logits, fp16 | 150 x 512 x 512 | 79 MB | 1.26 GB |
| Argmax label map, uint8 | 512 x 512 | 262 KB | 4.2 MB |
| RLE of a typical scene | variable | 5 to 40 KB | under 1 MB |
The table makes two decisions for you. Never copy logits to the host: moving 157 MB per batch across PCIe at a realistic 20 GB/s costs about 8 ms by itself, and the host then has to do the argmax on the CPU. And do not materialise full-resolution fp32 logits for a whole batch: 2.5 GB of transient memory per batch competes with the model and with other tenants on the GPU. Upsample in fp16, upsample in chunks of images, or for models where it does not cost accuracy, take the argmax at low resolution and upsample the label map with nearest-neighbour. Measure that last option against your validation set; bilinear upsampling of logits before argmax usually gives cleaner boundaries.
Post-processing and mask encoding
Post-processing on the GPU is a few lines of PyTorch. The important detail is the metadata: you must undo the padding before resizing back to the original size, or every mask drifts toward one corner.
import torch
import torch.nn.functional as F
def postprocess_semantic(logits, metas, chunk=4):
"""logits: [B, C, h, w] fp16 on GPU. metas: list of dicts recorded at preprocess time:
{"orig_h", "orig_w", "scaled_h", "scaled_w"} where the scaled image sat at the top-left
of a padded square canvas of side S. Returns a list of uint8 CPU masks."""
out = []
S = None
for start in range(0, logits.shape[0], chunk):
block = logits[start:start + chunk]
S = S or block.shape[-1] * 4 # model predicts at stride 4
up = F.interpolate(block, size=(S, S), mode="bilinear", align_corners=False)
labels = up.argmax(dim=1).to(torch.uint8) # [b, S, S], 1 byte per pixel
for i, m in enumerate(metas[start:start + chunk]):
crop = labels[i, : m["scaled_h"], : m["scaled_w"]] # undo padding first
crop = F.interpolate(crop[None, None].float(),
size=(m["orig_h"], m["orig_w"]),
mode="nearest")[0, 0].to(torch.uint8) # labels: nearest only
out.append(crop)
# sizes differ per image, so each mask is copied separately (uint8 keeps it cheap)
return [t.cpu() for t in out]Two rules are encoded there. Label maps are resized with nearest-neighbour only, because bilinear interpolation of class indices invents classes that sit numerically between two real ones. And the crop happens before the final resize, which is the exact inverse of the letterbox applied at preprocessing time.
Encoding is the next cost. Run-length encoding (RLE) in the COCO convention scans the mask in column-major order and stores alternating run lengths of zeros and ones, starting with zeros. The reference implementation is pycocotools, which expects a Fortran-ordered uint8 array:
import numpy as np
from pycocotools import mask as mask_utils
def encode_instances(binary_masks): # list of HxW uint8 arrays with values 0/1
rles = []
for m in binary_masks:
rle = mask_utils.encode(np.asfortranarray(m))
rle["counts"] = rle["counts"].decode("ascii") # bytes -> JSON-safe string
rles.append(rle)
return rlesFor semantic maps, a single-channel PNG of the label map is often smaller than one RLE per class and every client can decode it. Polygons suit annotation tools but are lossy, so never make them the canonical output. Run encoding in a CPU worker pool: it is where a well-tuned GPU pipeline quietly stalls.
Promptable models: split encoder and decoder
Promptable models change the shape of the service. The original Segment Anything model (SAM, 2023) has a large ViT image encoder that takes a 1024 by 1024 input and produces a 256-channel embedding on a 64 by 64 grid, plus a small prompt encoder and mask decoder that turn a click or box into masks in tens of milliseconds. SAM 2 (2024) moved to a hierarchical Hiera encoder and added a memory mechanism so a prompt on one video frame propagates to later frames. SAM 3 (late 2025) adds concept prompts, a short noun phrase or an image exemplar, and returns every matching instance. Check the model card of the exact checkpoint you deploy for input sizes and licence terms rather than relying on any summary, including this one.
The serving pattern that follows from this architecture is a split. Run the encoder once per image as its own batched model, store the embedding, and run the decoder once per prompt. An interactive annotation session issues dozens of clicks on the same image, so the split turns dozens of heavy encoder calls into one. The embedding for SAM's 64 x 64 x 256 grid is about 4 MB in fp32 or 2 MB in fp16, so a GPU-resident cache with a 4 GB budget holds roughly 2,000 images.
import hashlib
from collections import OrderedDict
class EmbeddingCache:
"""LRU cache of image embeddings on the GPU, bounded by bytes, not entries."""
def __init__(self, budget_bytes):
self.budget, self.used, self.items = budget_bytes, 0, OrderedDict()
@staticmethod
def key(image_bytes, model_rev):
# hash content + model revision: a new checkpoint must never reuse old embeddings
return hashlib.sha256(model_rev.encode() + image_bytes).hexdigest()
def get(self, k):
if k in self.items:
self.items.move_to_end(k)
return self.items[k]
return None
def put(self, k, emb):
size = emb.element_size() * emb.nelement()
while self.items and self.used + size > self.budget:
_, old = self.items.popitem(last=False)
self.used -= old.element_size() * old.nelement()
self.items[k] = emb
self.used += size
def segment(image_bytes, prompts, cache, encoder, decoder, model_rev):
k = cache.key(image_bytes, model_rev)
emb = cache.get(k)
if emb is None:
emb = encoder(image_bytes) # heavy: batched separately, fp16
cache.put(k, emb)
return decoder(emb, prompts) # light: one call per click or boxRoute every request for a session to the same replica with a consistent hash on the image key, or each click pays the encoder cost again. Hash content, not a client-supplied id, so an edited image is never served a stale embedding.
Shapes, engines and batching
Segmentation inputs arrive in every aspect ratio, and engines are fastest with fixed shapes. Letterboxing everything to one square is simplest but wastes compute on padding. Bucketing by aspect ratio (square, 4:3, 16:9), each with its own TensorRT optimisation profile, wastes less but means more engines to test. Dynamic-shape engines are flexible, but kernels are tuned for the optimal shape. Build engines with explicit ranges so behaviour is predictable:
trtexec --onnx=segformer_b2.onnx --fp16 \
--minShapes=pixel_values:1x3x512x512 \
--optShapes=pixel_values:8x3x512x512 \
--maxShapes=pixel_values:16x3x512x512 \
--saveEngine=segformer_b2_512.planIn Triton, enable dynamic batching on the model and keep the queue delay short, because segmentation latency is already dominated by post-processing:
dynamic_batching {
preferred_batch_size: [ 4, 8 ]
max_queue_delay_microseconds: 2000
}
instance_group [ { count: 2, kind: KIND_GPU } ]Two instances let one batch run post-processing while the next runs the network. The general trade-offs between static, dynamic and continuous batching are covered in GPU Batching Strategies, and multi-model and ensemble deployment in Triton Inference Server.
Tiling images larger than the model
Satellite scenes, whole-slide pathology images and factory line-scan cameras are far larger than any model input. Downscaling destroys small objects, so you tile: cut the image into overlapping windows, segment each, and stitch. The overlap matters because predictions near a window edge lack context and are systematically worse. A standard approach keeps only the central region of each tile, or blends overlapping logits with a weight that falls off toward the edges:
def tiled_segment(image, model, tile=1024, overlap=128, num_classes=150):
H, W = image.shape[-2:]
acc = torch.zeros(num_classes, H, W, device="cuda", dtype=torch.float16)
wsum = torch.zeros(1, H, W, device="cuda", dtype=torch.float16)
r = torch.ones(tile, device="cuda", dtype=torch.float16)
edge = torch.linspace(0.05, 1.0, overlap, device="cuda", dtype=torch.float16)
r[:overlap], r[-overlap:] = edge, edge.flip(0) # linear ramp, floor 0.05 at borders
ramp = r[:, None] * r[None, :] # separable 2-D weight
step = tile - overlap
for y in range(0, max(H - overlap, 1), step):
for x in range(0, max(W - overlap, 1), step):
y0, x0 = min(y, max(H - tile, 0)), min(x, max(W - tile, 0))
patch = image[..., y0:y0 + tile, x0:x0 + tile]
logits = model(patch) # [C, tile, tile] after upsampling
acc[:, y0:y0 + tile, x0:x0 + tile] += logits * ramp
wsum[:, y0:y0 + tile, x0:x0 + tile] += ramp
return (acc / wsum.clamp_min(1e-3)).argmax(0).to(torch.uint8)In production, batch the tiles instead of looping, and watch memory: the accumulator for a 20,000 by 20,000 scene with 150 classes in fp16 is 120 GB. Accumulate per region and write finished regions out. Instance masks crossing a tile boundary need merging by IoU in the overlap.
Failure modes
- Shifted masks. The preprocess letterboxed to the centre but post-processing crops from the top-left, or the reverse. Offline evaluation that reuses the same buggy code agrees with itself. Test with a synthetic image containing a single square at a known position and assert the returned mask covers it exactly.
- Invented classes at boundaries. Label maps resized with bilinear interpolation. Use nearest-neighbour for labels, always.
- Label map drift. The model was retrained with a reordered class list and the server still uses the old index-to-name table. Ship the label map inside the model artefact and check its hash at load time.
- CPU-bound encoding. GPU utilisation is low, latency is high and CPU cores are saturated by PNG or RLE encoding. Size the encode pool, reduce mask count with a score threshold, and measure stages separately.
- Response blow-up. An instance model returns its maximum of 100 masks on a cluttered image and the response grows to megabytes. Cap instances per request, apply a score threshold server-side, and return a flag when the cap was hit.
- Memory spikes. Full-resolution fp32 logits for a large batch push the GPU into out-of-memory errors only on large inputs. Enforce a maximum input size at the gateway.
Operating it
Instrument each stage with its own histogram: preprocess, queue wait, network, post-process, encode and response bytes. A single end-to-end latency number hides the stage that is actually slow. Use GPU profiling to confirm whether interpolate and argmax kernels, not the network, dominate GPU time; they often do for small models. If several small segmentation models share a large GPU, isolate them with MIG partitions so one tenant's tiling job cannot starve another's interactive traffic.
For quality, run a golden set with ground-truth masks through the production path on every deploy and track IoU per class, since one small class can regress inside a healthy average. Canary on IoU between old and new masks for mirrored traffic, sliced by condition such as night images or camera model.
Trade-offs
| Choice | Gains | Costs |
|---|---|---|
| Argmax on GPU, uint8 to host | Tiny transfers, low CPU | Logits unavailable to clients |
| Low-res argmax, nearest upsample | Much less memory | Blockier boundaries; validate mIoU |
| Letterbox to one shape | One engine, simple batching | Wasted compute on padding |
| Aspect buckets | Less padding | More engines and profiles to test |
| Encoder and decoder split | Interactive clicks in tens of ms | Cache memory, session routing |
What to do next
- Write down the output arithmetic for your model, input size and batch, and decide where argmax happens before writing any serving code.
- Record resize, pad and crop metadata at preprocessing and add the single-square alignment test to CI.
- Move upsampling and argmax onto the GPU, in fp16 and in chunks, and copy only uint8 masks to the host.
- Pick one canonical mask encoding, size a CPU pool for it, and measure the encode stage on its own.
- For promptable models, split encoder and decoder, cache embeddings by content hash plus model revision, and add session affinity.
- Build engines with explicit shape profiles and enable dynamic batching with a short queue delay.
- Run a golden-set mIoU check through the production path on every deploy and canary on mask IoU against the previous model.