FlashInfer is an open-source library of GPU kernels for LLM inference serving. Its core is attention over a paged KV cache, and around that it offers the other operations a serving engine runs every step: sampling, RMSNorm, rotary embeddings and a growing set of GEMM and mixture-of-experts kernels. The design is described in the paper "FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving", which won a best paper award at MLSys 2025, and the library is integrated into SGLang, vLLM and MLC-Engine.

Most engineers meet FlashInfer indirectly, as an attention backend flag in their serving engine. This article explains what sits behind that flag: how FlashInfer represents a batch of ragged, paged sequences, why its API splits every step into a plan and a run, how it balances work when one request has 100,000 tokens of context and the next has 50, and what goes wrong in production. It assumes you know the basics of FlashAttention and the paged KV cache.

Advertisement

Why serving attention is a different problem

Training attention works on dense, equal-length batches, and FlashAttention's tiling makes it compute-bound. Serving has three properties that break that picture. Batches are ragged: every request has a different context length. The KV cache is paged: a request's keys and values sit in fixed-size pages scattered through a pool, not in one contiguous tensor. And most steps are decode steps, where each request contributes one new query token that must attend over its entire cached history.

Decode is memory-bound. For one query token, a head reads every cached key and value once and does about one multiply-add per element read: a dot product with each key, a weighted sum of the values. With a Llama-3-8B-style model, meaning 32 layers, 8 KV heads, head dimension 128 and bf16, each token of context costs 2 x 32 x 8 x 128 x 2 bytes = 128 KiB of KV cache. A batch of 64 requests averaging 4,000 tokens holds about 31 GiB of KV, and every decode step has to stream all of it from HBM. The kernel's job is to keep that stream at close to peak bandwidth, whatever the mix of sequence lengths and page layouts.

The block-sparse page table

FlashInfer treats the KV cache of a whole batch as a block-sparse matrix: rows are requests, columns are pages in the global pool, and a non-zero block means that request owns that page. It stores this in a compressed sparse row form with three integer tensors. indptr has batch_size + 1 entries, and request i owns entries indptr[i] to indptr[i+1] of indices, which lists page IDs in order. last_page_len says how many slots of each request's final page are filled, from 1 to page_size.

A worked example with page_size 16 and three requests whose KV lengths are 37, 16 and 5 tokens: the first needs three pages (16 + 16 + 5), the second exactly one full page and the third one partial page.

import math, torch

def page_table(kv_lens, page_size, first_free_page=0):
    """Build FlashInfer's CSR-style page metadata from per-request KV lengths."""
    indptr, indices, last = [0], [], []
    nxt = first_free_page
    for n in kv_lens:
        pages = math.ceil(n / page_size)
        indices.extend(range(nxt, nxt + pages))   # a real allocator hands out scattered pages
        nxt += pages
        indptr.append(indptr[-1] + pages)
        last.append(n - (pages - 1) * page_size)   # 1..page_size, never 0
    as_t = lambda x: torch.tensor(x, dtype=torch.int32, device="cuda")
    return as_t(indptr), as_t(indices), as_t(last)

indptr, indices, last_page_len = page_table([37, 16, 5], page_size=16)
# indptr = [0, 3, 4, 5]   indices = [0, 1, 2, 3, 4]   last_page_len = [5, 16, 5]

Each layer's cache is a tensor of shape [max_num_pages, 2, page_size, num_kv_heads, head_dim] in the NHD layout, where the 2 holds K and V. With the 8B-like shapes above, one page for one layer is 2 x 16 x 8 x 128 x 2 bytes = 64 KiB. The same format also expresses things that are not paging at all: a shared prefix is several rows pointing at the same pages, and page_size 1 turns it into a token-level sparse layout for tree attention in speculative decoding.

Advertisement

Plan and run: the inspector-executor split

Every FlashInfer batch wrapper has two calls. plan() takes the page-table metadata and the head configuration, inspects the shape of the batch, decides how to partition the work across thread blocks, and writes that schedule into a workspace buffer. run() takes the query tensor and one layer's KV cache and executes the attention using the schedule. The batch shape is identical in every layer of a step, so a serving engine calls plan once per step and run once per layer.

One decode step with FlashInfer: plan once on the host, run once per layerSchedulerwhich requests runPage tableindptr, indices, last_lenwrapper.plan()split-KV work partitionWorkspace buffertiles, partial stateswrapper.run() x Lattention per layerMerge partialsLSE-weighted combinePaged KV cache[pages, 2, 16, Hkv, D]Sampling kerneltop-k / top-pmetadatawrites planreusedKV pageslogitsplan() depends only on batch shape, so it runs once per step and is shared by all L layers.With CUDA graphs, buffers are pre-allocated and plan() refills them in place before replay.
The plan is computed from batch metadata before the first layer and reused by every layer. Split-KV partial results are merged before the output returns.
import torch, flashinfer

L, Hq, Hkv, D, P = 32, 32, 8, 128, 16          # Llama-3-8B-like shapes, page_size 16
max_pages = 4096
workspace = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda")
kv_cache = [torch.empty(max_pages, 2, P, Hkv, D, dtype=torch.bfloat16, device="cuda")
            for _ in range(L)]                  # NHD layout: [pages, K/V, page, heads, dim]

decode = flashinfer.BatchDecodeWithPagedKVCacheWrapper(
    workspace, "NHD",
    use_tensor_cores=True,                      # GQA group of 4 query heads per KV head
)

def decode_step(q_per_layer, indptr, indices, last_page_len):
    # Inspector: once per step, from batch metadata only
    decode.plan(indptr, indices, last_page_len, Hq, Hkv, D, P,
                q_data_type=torch.bfloat16, kv_data_type=torch.bfloat16)   # default is fp16
    outs = []
    for layer in range(L):
        # Executor: once per layer, reusing the same plan
        outs.append(decode.run(q_per_layer[layer], kv_cache[layer]))   # [batch, Hq, D]
    return outs

The split exists for two reasons. Scheduling costs host time, and paying it once instead of 32 times matters when a decode step takes a few milliseconds. And it makes the kernels compatible with CUDA graphs. A graph captures fixed kernel launches with fixed pointers, but a serving batch changes every step. With use_cuda_graph=True the wrapper uses pre-allocated index buffers that you pass in, plan refills them in place, and the captured graph replays with the new schedule. Engines typically capture one graph per padded batch size.

Load balancing with split-KV

Imagine a batch of 63 short requests and one request with 100,000 tokens of context. If each request maps to one thread block per head, the GPU finishes the short ones almost at once and then waits on a handful of blocks crawling through the long one, leaving most streaming multiprocessors idle. FlashInfer's scheduler instead splits long KV ranges into chunks, spreads chunks from all requests across the available blocks so each does roughly equal work, and computes a partial attention output for each chunk.

Partial results merge exactly, because softmax attention over a union of disjoint key sets can be recombined from each part's output and its log-sum-exp (LSE):

import torch

def merge_two(o1, lse1, o2, lse2):
    """Exact combination of attention over two disjoint KV chunks.
    o*: [batch, heads, D] partial outputs, lse*: [batch, heads] natural-log sum-exp."""
    m = torch.maximum(lse1, lse2)
    w1, w2 = torch.exp(lse1 - m), torch.exp(lse2 - m)
    o = (o1 * w1[..., None] + o2 * w2[..., None]) / (w1 + w2)[..., None]
    return o, m + torch.log(w1 + w2)

FlashInfer implements this merge as a kernel and exposes it as merge_state and merge_states, so you can combine attention computed by different kernels, or on different GPUs, yourself. If you mix FlashInfer's returned LSE with your own code, check which logarithm base its documentation specifies before combining. The cost of splitting is the extra write and read of partial outputs through the workspace, which is why the scheduler does not split short sequences, and why plan() must see real lengths to choose well.

GQA, tensor cores and MLA

Grouped-query attention changes decode's arithmetic. In the example, 32 query heads share 8 KV heads, so each KV element loaded feeds four query heads. Treating the four queries of a group as a small matrix turns decode into a thin matrix multiply that tensor cores can run, which raises effective throughput without reading any more memory. The decode wrapper exposes this as use_tensor_cores; it pays off as the group size grows, and the prefill kernels already work this way because they have many query tokens.

Multi-head latent attention, used by DeepSeek-style models, caches a compressed latent instead of per-head keys and values, and has its own wrapper, BatchMLAPagedAttentionWrapper. The page-table idea is the same, but the head dimensions and the absorbed projections differ, so do not try to route MLA through the standard wrappers.

Prefill, masks and shared prefixes

BatchPrefillWithPagedKVCacheWrapper handles the case with several query tokens per request: prompt prefill, chunked prefill and verification of speculative drafts. Its plan takes a qo_indptr describing ragged query lengths alongside the page table, and options including causal, window_left for sliding-window attention, logits_soft_cap for models that cap attention logits, and a custom mask. For tree-structured drafts, see how verification works in speculative decoding.

When many requests share a long prefix, such as a system prompt or a few-shot block, MultiLevelCascadeAttentionWrapper computes attention over the shared pages once for all queries, computes each request's unique suffix separately, and merges the two with the same LSE combine. That turns many reads of the prefix into one, which pays off when the shared part is long relative to the suffixes.

JIT, variants and backends

Attention in real models is not one formula. Soft capping, sliding windows, ALiBi-style biases, custom masks and different head dimensions each change the inner loop. FlashInfer generates kernels from templates, compiling a variant just-in-time the first time a combination is requested and caching it. The library also exposes a backend choice: the current source lists auto, fa2, fa3, cudnn, trtllm-gen and CuTe-DSL based options, and auto picks based on GPU architecture and problem shape.

The practical consequence is cold-start cost. The first request that hits a new variant can wait for compilation, and a fleet that autoscales onto fresh nodes pays it on every node unless the cache is prebuilt or persisted. Warm up every shape and variant you serve before a node takes traffic, exactly as you warm up CUDA graphs.

Sampling and the rest of the step

After attention and the rest of the layers comes sampling, which on a 128,000-token vocabulary is not free. Naive top-p sorts the whole vocabulary per request. FlashInfer's sampling kernels, such as top_k_top_p_sampling_from_logits and min_p_sampling_from_probs, use rounds of rejection sampling inside one kernel instead of a sort, and chain_speculative_sampling implements the accept and reject step for speculative decoding. Fused kernels such as fused_add_rmsnorm and the apply_rope family cut the small, memory-bound operations between the big GEMMs. These kernels save little each, but a decode step is many small operations and their launch and memory overheads add up.

Failure modes in production

  • Stale plan. Changing the batch, for example admitting or evicting a request, without calling plan again means run uses a schedule for the wrong lengths. Symptoms range from garbage output to illegal memory access. Tie plan to the scheduler's batch commit.
  • Off-by-one page lengths. last_page_len must be between 1 and page_size. Writing 0 for a request whose length is a multiple of the page size, instead of page_size with one fewer page, silently drops or reads 16 tokens.
  • Dtype mismatch between plan and run. Plan specialises on query and KV dtypes. Planning for fp16 and running with bf16 or fp8 KV either fails or picks the wrong kernel. Pass the dtypes explicitly.
  • Workspace too small. Large batches with heavy splitting need more workspace for partial outputs. Size it for your worst-case batch, not the average.
  • JIT cold start under load. Latency spikes on the first request of a new shape, often right after a scale-out. Pre-warm.
  • Version skew. Serving engines pin a FlashInfer version and call internal arguments that change between releases. Upgrade the library only together with the engine, and re-run accuracy checks after changing backend.

Trade-offs and when to reach for it

If you run vLLM or SGLang, FlashInfer is one of several attention backends and the engine owns the integration; your decisions are which backend to select for your GPU and model, and how to benchmark it on your real traffic mix. vLLM on GPU covers where that fits in the engine. Call FlashInfer directly when you are building your own engine, research kernels or unusual attention patterns, and budget for owning the page allocator, the plan lifecycle and graph capture yourself.

The benefits are largest where serving is irregular: highly variable context lengths, long shared prefixes and speculative trees. For uniform offline batches of short prompts, a dense FlashAttention kernel may be just as fast and simpler to operate.

What to do next

  1. Check which attention backend your serving engine uses today and benchmark FlashInfer against it on a replay of real request lengths, not a uniform synthetic batch.
  2. If you integrate directly, write the page-table builder first and unit-test it on lengths of 1, page_size, page_size + 1 and exact multiples.
  3. Call plan once per step from the batch commit, run once per layer, and assert that plan was called for the current batch ID.
  4. Enable tensor-core decode for GQA models and compare throughput at your real group size.
  5. Pre-warm JIT variants and CUDA graphs for every batch size and model shape before a node takes traffic.
  6. Pin FlashInfer together with your engine version and gate upgrades on accuracy and latency regression tests.
Key takeaway: FlashInfer makes serving attention fast by treating a batch of ragged, paged sequences as one block-sparse matrix. It inspects the batch once per step in plan() and reuses that schedule across every layer in run(), which also keeps it compatible with CUDA graphs. It balances long and short requests by splitting KV ranges and merging partial results exactly through their log-sum-exp. Get the page metadata right, re-plan whenever the batch changes, pre-warm JIT variants and graphs, and move the library version only in lockstep with your engine.