Serving a Vision Transformer looks easy on a slide: one forward pass, no decode loop, no KV cache. In production it hides three traps. Cost depends on resolution in a way that turns quadratic, so a ViT-B/16 that needs 17.6 billion multiply-accumulates at 224 pixels needs about 37 times as many at 1024. Small batches leave a large GPU idle, so throughput depends on how well you group requests of the same shape. And the shape itself is a choice you make in preprocessing, which silently changes accuracy if training and serving disagree.

This article treats the ViT as a standalone service: an image classifier, an embedding endpoint for retrieval, or a dense model such as a segmentation backbone. It builds the cost model from first principles, shows where attention starts to dominate, and then walks through the serving stack: resolution buckets, a batcher, captured CUDA graphs, fused attention, token merging and quantization. If you serve the vision tower inside a VLM, read the vision encoder article and the multimodal serving article alongside this one; for JPEG decode on the GPU and TensorRT engine building, see image classification serving.

The cost model: where attention takes over

A ViT cuts the image into P by P patches, projects each to width d, prepends a class token, and runs L identical blocks. Each block does four d-by-d projections for attention (Q, K, V and output) and an MLP that expands to 4d and back: twelve d-squared multiply-accumulates (MACs) per token. Attention itself adds two matmuls, QK-transpose and the weighted sum over V, each n by n by d for n tokens. Per image:

n    = (res // P) ** 2 + 1              # patches plus the class token
MACs = L * (12 * d * d * n + 2 * n * n * d) + patch_embed
# attention share of block MACs = 2 n^2 d / (12 d^2 n + 2 n^2 d) = n / (6d + n)

For ViT-B/16 (d = 768, L = 12) at 224 pixels, n = 197 and the total is about 17.6 billion MACs, the figure usually quoted as 17.6 GFLOPs. The attention matmuls are 4.1 percent of it, so the model behaves like a stack of dense layers. The share is n / (6d + n): it reaches one half when n = 6d, which for ViT-B is 4,608 tokens, roughly a 1,088-pixel image at patch 16. The chart below shows the same curve for ViT-L/16 (d = 1024, L = 24).

Two consequences follow. First, cost per image grows faster than pixel count once n is a few thousand, so price and capacity plans must be made per resolution, never per image. Second, memory changes character: an unfused attention kernel materialises an n by n score matrix per head. For ViT-L at 1024 pixels that is 16 heads times 4,097 squared times two bytes, about 537 MB per image per layer in FP16, which is why fused attention is not an optimisation but a requirement at high resolution.

Fused attention is mandatory at high resolution

PyTorch exposes fused attention through torch.nn.functional.scaled_dot_product_attention (SDPA), which dispatches to a FlashAttention-style kernel, a memory-efficient kernel or a plain math fallback depending on dtype, head size, mask and hardware. The fused kernels tile the computation so the score matrix lives in on-chip SRAM and never reaches HBM; see FlashAttention for the mechanism. Modern timm and Hugging Face ViT implementations call SDPA already, but a custom block, an attention mask of the wrong dtype or FP32 inputs can push dispatch onto the math path without any error. Pin the backend in a test so a silent fallback fails loudly:

import torch
from torch.nn.attention import sdpa_kernel, SDPBackend

@torch.inference_mode()
def assert_fused_attention(model, res, batch=8):
    x = torch.randn(batch, 3, res, res, device="cuda", dtype=torch.bfloat16)
    # Raises if no allowed backend can run this shape, instead of falling back to math.
    with sdpa_kernel([SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]):
        model(x)

Dense-prediction backbones take a different route. ViTDet-style models, including the image encoder of Segment Anything, process 1024-pixel inputs with 14 by 14 windowed attention in most blocks and a few global attention blocks, which keeps the quadratic term bounded. If you serve such a model, profile the global blocks separately: they dominate the tail when resolution grows.

Resolution buckets and CUDA graphs

A ViT is trained with learned position embeddings for one grid size. Serving at another size means interpolating those embeddings, which works but shifts accuracy, so the first decision is a small set of supported resolutions, called buckets. Every request is resized or padded to the nearest bucket edge, and each bucket gets its own queue. Recent timm versions can build a ViT with dynamic_img_size=True to interpolate on the fly; evaluate accuracy per bucket before you enable a size the model never saw in training.

Buckets make shapes static, and static shapes unlock CUDA graphs. A ViT forward pass at batch 8 launches hundreds of small kernels; at low resolution the CPU launch overhead can exceed GPU time. Capturing the forward once per (bucket, batch size) pair and replaying it removes that overhead. torch.compile(model, mode="reduce-overhead") captures graphs for you; the manual version makes the mechanics visible:

import torch

class GraphedViT:
    """One captured CUDA graph per (resolution, batch) pair; inputs padded to fit."""

    def __init__(self, model, buckets=(224, 384), batches=(1, 4, 8, 16, 32)):
        self.model, self.graphs = model.eval(), {}
        for res in buckets:
            for b in batches:
                x = torch.zeros(b, 3, res, res, device="cuda", dtype=torch.bfloat16)
                s = torch.cuda.Stream()
                s.wait_stream(torch.cuda.current_stream())
                with torch.cuda.stream(s), torch.inference_mode():
                    for _ in range(3):          # warm up allocator and autotuning
                        self.model(x)
                torch.cuda.current_stream().wait_stream(s)
                g = torch.cuda.CUDAGraph()
                with torch.cuda.graph(g), torch.inference_mode():
                    y = self.model(x)
                self.graphs[(res, b)] = (g, x, y)

    def __call__(self, images, res):
        n = images.shape[0]
        b = min(k for (r, k) in self.graphs if r == res and k >= n)
        g, x, y = self.graphs[(res, b)]
        x[:n].copy_(images)                     # static input buffer
        x[n:].zero_()                           # padding rows are wasted work
        g.replay()
        return y[:n].clone()                    # copy out before the next replay

The batch ladder trades padding waste against graph count. With rungs 1, 4, 8, 16 and 32, a batch of 9 runs as 16 and wastes 44 percent of that replay; adding a 12 rung cuts it. Track padding waste as a metric per bucket and add rungs where the histogram says requests land. Each graph pins its own activation memory, so count graphs times workspace against the memory budget before adding many.

Batching per bucket: a worked example

Batching follows the classic rule: wait until the batch is full or the oldest request has waited T milliseconds, whichever comes first. What is ViT-specific is that the right B and T differ by bucket, because GPU time per image differs by up to an order of magnitude. Use a worked example. Suppose profiling ViT-B/16 on one GPU gives 6 ms for a batch of 32 at 224 pixels and 20 ms for a batch of 32 at 384 pixels (measure your own; these are illustrative). The p99 target is 50 ms end to end.

  • At 224: a full batch adds at most 6 ms of compute. If decode and network take 15 ms, you can afford T up to about 25 ms and still meet the target with headroom for one queued batch ahead of you.
  • At 384: compute is 20 ms, so one batch ahead plus your own already costs 40 ms. T must shrink to a few milliseconds, and B should drop to 16 so that a queued batch finishes sooner.
  • Capacity: at 224, 32 images per 6 ms is about 5,300 images per second if the queue is always full; at 384, 32 per 20 ms is 1,600. Plan replicas from the traffic mix per bucket, not from a single average.

Keep the batcher in the same process as the GPU worker, so there is no extra network hop, and give each bucket its own CUDA stream only if profiling shows the GPU is idle between batches. Two concurrent high-resolution batches usually just contend for the same SMs.

ViT-L/16 multiply-accumulates per image by resolution224 px197 tokens61.6 G384 px577 tokens191.1 G512 px1025 tokens362.0 G1024 px4097 tokens2065.5 Glinear layers (QKV, proj, MLP)attention matmuls (QK^T, AV)
MACs per image for ViT-L/16 by resolution; the red part is the attention matmuls, which grow with the square of the token count.

Fewer tokens: merging and resolution choice

The cheapest token is one you do not process. Token Merging (ToMe, Bolya et al., 2022) merges similar tokens between attention and MLP in each block using a fast bipartite matching, removing r tokens per block. Its paper reports about twice the throughput of ViT-L at 512 pixels and ViT-H at 518 pixels with a 0.2 to 0.3 percent accuracy drop, without retraining. The reference implementation patches timm models in two lines:

import timm, tome

model = timm.create_model("vit_base_patch16_224", pretrained=True)
tome.patch.timm(model)
model.r = 16        # tokens merged per block; sweep it against your eval set

Three cautions. Merging changes the token count per block, so verify that your CUDA graphs and any feature taps still see the shapes they expect. Merged tokens break the one-to-one map between tokens and patches, so dense tasks that need a per-patch output must unmerge, or skip ToMe. And accuracy loss is task-specific: measure it on your data, including the rare classes, not just top-1 on a benchmark.

The other lever is resolution itself. Many retrieval and classification workloads lose little accuracy when served at 224 instead of 384, at about a third of the compute, while OCR-like and small-object tasks lose a lot. Run the same eval at each candidate bucket and choose per product surface.

Quantizing a ViT without losing the plot

ViTs are harder to quantize than CNNs. Post-softmax attention probabilities and post-GELU activations have skewed distributions, and LayerNorm inputs vary strongly across channels; papers such as PTQ4ViT and FQ-ViT exist precisely because naive per-tensor INT8 loses noticeable accuracy. A practical order of operations:

  1. Serve in BF16 first and record accuracy per bucket. This is the baseline every later change is compared with.
  2. Quantize only the linear layers (QKV, projection, MLP) to INT8 or FP8 with per-channel weight scales, keep softmax, LayerNorm and GELU in higher precision, and calibrate activation scales on a few hundred real production images, not on ImageNet.
  3. Compare logits or embeddings, not just accuracy: compute cosine similarity between BF16 and quantized embeddings on a held-out set. For a retrieval index, a drop in similarity means you must re-embed the corpus, because old and new vectors no longer live in the same space.
  4. Only then try more aggressive schemes, and keep the BF16 path deployable as a rollback.

Running an embedding endpoint

Embedding endpoints have operational rules of their own. The vector is a contract: every stored vector was produced by a specific checkpoint, preprocessing recipe and precision, and a query embedded under a different combination searches the wrong space. Put the model revision, bucket and normalisation in the response and in the index metadata, and refuse to query an index with a mismatched revision.

Cache by content. Hash the image bytes together with the preprocessing configuration and model revision; repeated product photos and re-uploaded documents are common, and a cache hit skips decode as well as the forward pass. Normalise vectors once, at the server, so clients cannot disagree about it.

Standalone ViT serving: request to responseClientimage bytes + taskAdmission + hashsize limit, cache keymissDecode + resizeto bucket edgeBucket queue224 / 384 / 512Batchermax batch B or deadline TCaptured graphone per (bucket, batch)ViT forward on the GPUPatch embedconv 16x16L blocksSDPA + MLPHeadCLS / pool / mapOutputlogits, vector, maskhitMetrics per bucketqueue wait, batch fill, GPU ms, padding waste, cache hit rate, p99
Admission and content hashing come first; misses are decoded straight to a bucket edge, batched per bucket and run through a pre-captured graph.

Failure modes

FailureSymptomFix
Train/serve resize mismatchAccuracy drops a few points with no errorsUse the model's own preprocessing config (interpolation, crop ratio, mean/std); test by comparing logits on fixed images against the training pipeline
Attention falls back to the math kernelMemory spikes and latency triples at high resolutionPin SDPA backends in a CI test; keep inputs in BF16/FP16 and masks in a supported dtype
One giant image starves a bucketp99 spikes while p50 is flatCap input pixels at admission and route oversized inputs to a separate queue or tile them
Graph replay reads stale inputOutputs from the previous batchCopy into the static buffer before replay and clone outputs before the next one
Padding wasteGPU busy, throughput lowTrack padded rows per replay; add batch rungs where requests cluster
Silent embedding driftRetrieval recall drops after a deployVersion vectors; block queries across revisions; re-embed on any model or precision change
CPU decode bottleneckGPU utilisation under 50 percent with a full queueDecode on the GPU or add decode workers; measure decode ms separately

Trade-offs

ChoiceGainsCosts
Lower serving resolutionRoughly linear-to-quadratic compute savingSmall-object and text accuracy
More bucketsLess resize distortion and paddingMore graphs, memory and queues to tune
Larger batch or longer waitThroughputTail latency, especially at high resolution
Token mergingUp to about 2x at high resolutionTask-specific accuracy loss; breaks dense outputs
INT8 or FP8 linearsMemory and throughputCalibration work; embedding-space drift
torch.compile vs TensorRTCompile: easy iteration; TensorRT: often the fastest engineCompile: warm-up time; TensorRT: rebuild per shape and version

What to do next

  1. Compute n and the attention share n / (6d + n) for your model at every resolution you serve, and write the per-bucket cost table into the capacity plan.
  2. Run the SDPA backend test above in CI for each bucket.
  3. Choose two or three buckets from the real input-size histogram, capture graphs per (bucket, batch rung) and record padding waste.
  4. Measure GPU ms per batch for each bucket, then set B and T per bucket from the latency target using the arithmetic in the worked example.
  5. Evaluate 224 versus higher buckets and ToMe settings on your own eval set before shipping either.
  6. Version every embedding with model revision, precision and preprocessing, and add a content-hash cache in front of the GPU.
  7. Read CLIP training on GPU if you also train the encoders you serve.
Key takeaway: A ViT's cost is set by its token count, and attention's share grows as n / (6d + n), so plan capacity per resolution. Make shapes static with a few resolution buckets, capture a CUDA graph per bucket and batch rung, set batch size and wait time per bucket from measured GPU time, insist on fused attention, and treat token merging, lower resolution and quantization as accuracy trades to measure on your own data. Version every embedding you emit.