PidokuInfra

Static vs Dynamic Shapes

Basic Advanced 1h Difficulty 3/5 Topic 11 of 12

Prerequisites 08, 10


1. What is it?#

Static shapes are known at compile time. Dynamic shapes vary per request. LLM serving is inherently dynamic — prompt lengths, batch sizes, and generation lengths all vary — while nearly every optimization technique wants static.

This tension is the central engineering compromise of production inference.


2. Why does it exist?#

Optimization             wants
─────────────────────────────────────────
CUDA graphs              static shapes AND static addresses
AOT compilation          static shapes
Kernel autotuning        known shapes
Memory pre-planning      known shapes
Tensor core efficiency   dimensions that are nice multiples

Reality                  gives
─────────────────────────────────────────
Prompt length            7 to 128,000 tokens
Batch size               1 to 256, changing every step
Generation length        1 to 8,000 tokens, unknown in advance

Something has to give. What gives is: bucketing, padding, and accepting some waste.


3. Simple analogy#

Shoe sizes. Feet are continuous; manufacturing wants discrete. So you make sizes 6, 7, 8, 9, and everyone wears the nearest one up. Slight waste (a size-8.2 foot in a size-9 shoe), huge manufacturing efficiency.

Bucketing shapes is exactly this, and the design question is exactly the same: how many sizes, and where do you put them?


4. Tiny example#

The cost of padding, made concrete:

Python
import torch, time

d = 4096
W = torch.randn(d, 4*d, device='cuda', dtype=torch.float16)

def t(M):
    A = torch.randn(M, d, device='cuda', dtype=torch.float16)
    for _ in range(10): A @ W
    torch.cuda.synchronize(); t0 = time.perf_counter()
    for _ in range(50): A @ W
    torch.cuda.synchronize()
    return (time.perf_counter()-t0)/50

for M in [8, 9, 16, 17, 32, 33, 64, 65, 128, 129]:
    print(f"M={M:4d}  {t(M)*1e6:8.1f} us")

You’ll see steps: M=9 costs almost exactly what M=16 costs. So if your graph-capture buckets are {1,2,4,8,16,32,…}, running batch 9 padded to 16 wastes nothing in that GEMM — the hardware was going to round up anyway.

But it does waste KV-cache reads and attention work, which scale with the real batch. So the padding cost is real but smaller than naive analysis suggests.


5. Technical explanation#

The four dynamic dimensions#

1. BATCH SIZE (number of sequences)
   Varies every decode step as requests join and leave.
   Handled by: CUDA graph buckets + padding

2. PROMPT LENGTH (prefill sequence length)
   7 to 128k. Enormous variance.
   Handled by: varlen kernels (cu_seqlens), chunked prefill

3. CONTEXT LENGTH (KV cache size per sequence)
   Grows every step.
   Handled by: paged KV cache + block tables (the kernel takes the length as data)

4. NUMBER OF PREFILL vs DECODE tokens in a step
   Chunked prefill mixes them.
   Handled by: unified kernels that take per-request metadata

Strategy 1: bucketing + padding#

Capture CUDA graphs for batch ∈ {1,2,4,8,16,24,32,48,64,96,128,192,256}
Actual batch 37 → pad to 48 → waste 23% of KV reads and attention work
                            → GEMM cost was going to round up anyway

Tradeoff: more buckets = less padding waste, but more capture time and memory. vLLM’s default capture list is configurable; tuning it for your batch-size distribution is a real (if minor) optimization.

Strategy 2: varlen kernels#

Instead of padding sequences to a common length, pack them and pass boundary offsets:

Padded:                          Packed (varlen):
[t t t t P P P P]                tokens: [t t t t t t t t t t t t]
[t t P P P P P P]                cu_seqlens: [0, 4, 6, 12]
[t t t t t t P P]
 12 real, 12 padding              12 real, 0 padding

FlashAttention’s flash_attn_varlen_func takes exactly this. For a realistic length distribution (mean 300, max 8000), padding wastes 60-95% of prefill compute; varlen wastes 0. This is not a micro-optimization; it’s essential.

Strategy 3: shape-agnostic kernel design#

Some kernels take the varying dimension as a runtime argument rather than a compile-time constant:

Attention over the KV cache: the sequence length is data (from the block table),
not a template parameter. The kernel loops until it runs out of blocks.
→ no recompilation, no bucketing needed for context length

This is why paged attention handles arbitrary context lengths without recompiling. The paged design solves a compilation problem as well as a memory problem — a point often missed.

Strategy 4: dynamic shape support in compilers#

Python
torch.compile(model, dynamic=True)     # generate kernels with symbolic shapes

The compiler marks dimensions as symbolic and generates code that works for any value. Costs:

  • Loses some optimizations (can’t specialize tile sizes to the exact shape).
  • Guards must be checked at runtime.
  • Typically 5-15% slower than fully static, but avoids recompilation storms.

PyTorch’s default is “compile static on first call, recompile dynamic if the shape changes” — usually the right heuristic.


6. Under the hood#

Watch recompilation happen:

Shell
TORCH_LOGS="recompiles,dynamic" python script.py

You’ll see messages like:

Recompiling function forward
    triggered by the following guard failure(s):
    - tensor 'x' size mismatch at index 0. expected 8, actual 9

After cache_size_limit (default 8) recompilations, Dynamo gives up:

torch._dynamo hit config.cache_size_limit (8)
   function: 'forward'
   → falling back to eager

Silent permanent performance loss. Monitor for this message in production logs.


7. Performance implications#

Approach                       Perf vs ideal static   Flexibility
Fully static (TensorRT, fixed) 100%                   none
Bucketed + padded              85-97%                 good
Dynamic-compiled               85-95%                 excellent
Varlen kernels                 ~100% (for prefill)    excellent
Eager (no compilation)         60-75%                 total

The practical answer for LLM serving: varlen kernels for prefill, bucketed CUDA graphs for decode, paged attention for context length. That’s what production engines do, and it gets you 90-95% of static performance with full flexibility.


8. Production implications#

  • Tune your CUDA graph capture list to your batch-size distribution. Measure the distribution first.
  • Ensure varlen kernels are used for prefill. Padding waste is the largest and most commonly overlooked inefficiency in naive serving code.
  • Cap max input length at the gateway. Unbounded prompt length means unbounded activation memory and unbounded prefill time.
  • Alarm on recompilation. A recompilation storm is a latency incident.
  • For TensorRT-LLM, choose your build profile carefully. max_batch_size, max_input_len, max_output_len, and max_num_tokens are baked in; too small and requests fail, too large and you waste memory.

9. Common mistakes#

Padding all sequences to the maximum. 60-95% waste on realistic distributions.

Too few CUDA graph buckets. Excessive rounding-up waste.

Not capping input length. One 500k-token request OOMs the server.

Recompilation storms from unbounded shape variation.

Assuming dynamic shapes are free. They cost 5-15%.

Building a TensorRT engine with max_input_len smaller than production traffic. Requests fail hard.


10. Hands-on exercise#

A. Padding waste. Sample 1,000 sequence lengths from a realistic distribution (lognormal, µ=5.5, σ=1.2, clipped at 8192). Compute total tokens vs total padded tokens for batch sizes 8, 32, 128. What fraction is waste?

B. Bucket design. Given a measured batch-size distribution, choose a set of 10 CUDA graph buckets that minimizes expected padding waste. Compare to a naive powers-of-two set.

C. Varlen. Implement attention with padding and with flash_attn_varlen_func for a batch of variable-length sequences. Compare time and correctness.

D. Recompilation. Compile a model with dynamic=False, then call it with 10 different shapes. Count recompilations and total stall time. Repeat with dynamic=True.

E. TensorRT profile. If you have TensorRT-LLM, build engines with two different max_num_tokens settings and compare memory and throughput.


11. Interview questions#

  1. Why do LLM workloads have dynamic shapes, and which dimensions vary?
  2. What is bucketing, and how do you choose the buckets?
  3. What are varlen kernels and how much do they save?
  4. How does paged attention avoid the dynamic-shape problem for context length?
  5. What is a recompilation storm and how do you detect it?
  6. What does dynamic=True cost you in torch.compile?
  7. Why must you cap max input length at the gateway?

12. Further reading#

↑↓ navigate↵ openesc close