FlashInfer Attention Kernels

Experimental Task-Scheduled Attention

The experimental Blackwell task-scheduled FMHA context, FMHA decode, block-sparse FMHA, and MLA decode APIs are imported from flashinfer.attention.prims_ts. Scheduling, tile selection, and split-KV reduction are automatic implementation details; there are no public tuning knobs.

See the PrimTS guide index for the public entry points, supported contracts, and examples. Current accuracy and performance signoff is on SM100a/B200; SM103a/B300 is architecture-gated but not yet signoff-qualified.

FMHA Context/Prefill

batch_prefill(q, k, v, *[, qo_indptr, ...])

Run one-shot fixed or packed-ragged task-scheduled context attention.

batch_prefill_with_paged_kv_cache(q, ...[, ...])

Run one-shot packed-Q context attention over separate HND page pools.

class flashinfer.attention.prims_ts.BatchPrefillTSWrapper

Compile and reuse fixed or packed-ragged contiguous context attention.

plan accepts only static compilation geometry. Q/K/V tensors, packed cumulative offsets, variable-window metadata, and optional scale overrides are supplied to run. Variable-window metadata may include precomputed CTA minima. Packed offset values may change between launches while their request lengths stay within the planned capacities.

Plan-time scalar defaults are stored as one-element device tensors. Variable-window plans also own private fallback reduction scratch refreshed when a run omits precomputed CTA minima; only launches using that scratch must not overlap across streams or captured graphs. With caller-owned output, validate=False performs no allocation or metadata readback and is CUDA-graph-capturable when the caller guarantees the complete runtime contract and stable addresses. Replanning replaces plan-owned tensors and invalidates graphs captured from the previous plan; all prior launches and replays must finish before plan is called again. Keep the wrapper and all captured runtime tensors alive until every graph using the current plan is destroyed. Before running on a CUDA stream that is not already ordered after the planning stream, the caller must establish that dependency.

__init__() → None

Initialize an unplanned task-scheduled context-attention wrapper.

plan(*, device: int | str | device, batch_size: int, max_seq_len_q: int, max_kv_len: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, q_dtype: dtype, kv_dtype: dtype, out_dtype: dtype | None = None, packed: bool = False, mask_type: Literal['dense', 'causal', 'variable_window'] = 'dense', window_left: int = -1, sm_scale: float | None = None, output_scale: float = 1.0) → None

Compile one reusable specialization from explicit static geometry.

Packed plans conservatively support any positive per-request Q and K/V lengths bounded by max_seq_len_q and max_kv_len. Fixed plans use those bounds as their exact tensor extents. batch_size is exact. Calling plan again replaces all plan-owned tensors and invalidates CUDA graphs captured from the previous plan. Complete every launch and replay that uses the previous plan before replanning. Planning allocates and compiles; complete it before CUDA Graph capture.

Parameters:
  • device (int, str, or torch.device) – CUDA device on which the specialization is compiled and run.

  • batch_size (int) – Exact number of requests in every run.

  • max_seq_len_q (int) – Maximum query length of any request.

  • max_kv_len (int) – Maximum K/V length of any request.

  • num_qo_heads (int) – Number of query/output heads.

  • num_kv_heads (int) – Number of key/value heads.

  • head_dim (int) – Query, key, value, and output head dimension.

  • q_dtype (torch.dtype) – Query dtype.

  • kv_dtype (torch.dtype) – Key/value dtype. It must currently equal q_dtype.

  • out_dtype (torch.dtype, optional) – Output dtype; defaults to q_dtype.

  • packed (bool) – Whether run tensors use packed [total_tokens, H, D] storage.

  • mask_type ({"dense", "causal", "variable_window"}) – Attention mask mode. variable_window is supported only for fixed-shape inputs.

  • window_left (int) – Left sliding-window extent, or -1 to disable the window.

  • sm_scale (float, optional) – Softmax scale; defaults to the inverse square root of head size.

  • output_scale (float) – Scale applied to the attention output.

run(q: Tensor, k: Tensor, v: Tensor, qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, *, variable_window_token_starts: Tensor | None = None, variable_window_token_ends: Tensor | None = None, variable_window_cta_starts: Tensor | None = None, out: Tensor | None = None, scale_softmax_log2: Tensor | None = None, output_scale: Tensor | None = None, validate: bool = True) → Tensor

Launch the compiled specialization with per-request tensors.

Parameters:
  • q (torch.Tensor) – Runtime query, key, and value tensors matching the plan.

  • k (torch.Tensor) – Runtime query, key, and value tensors matching the plan.

  • v (torch.Tensor) – Runtime query, key, and value tensors matching the plan.

  • qo_indptr (torch.Tensor, optional) – Cumulative packed Q and K/V offsets. Required by packed plans and forbidden by fixed plans.

  • kv_indptr (torch.Tensor, optional) – Cumulative packed Q and K/V offsets. Required by packed plans and forbidden by fixed plans.

  • variable_window_token_starts (torch.Tensor, optional) – Inclusive K-token bounds for every fixed Q row. Required only for mask_type='variable_window' and shaped [B, max_seq_len_q].

  • variable_window_token_ends (torch.Tensor, optional) – Inclusive K-token bounds for every fixed Q row. Required only for mask_type='variable_window' and shaped [B, max_seq_len_q].

  • variable_window_cta_starts (torch.Tensor, optional) – Precomputed exact minimum start bound for each Q work tile, shaped [B, ceil(max_seq_len_q / tile_size_q)] where tile_size_q is 256 for head dimension 128 and 128 for head dimension 256. Valid only for mask_type='variable_window'. When omitted, the wrapper derives these values from variable_window_token_starts.

  • out (torch.Tensor, optional) – Caller-owned output tensor. A new tensor is allocated when omitted.

  • scale_softmax_log2 (torch.Tensor, optional) – Per-run one-element float32 softmax scale in base-2 form. Defaults to the scale retained by the plan.

  • output_scale (torch.Tensor, optional) – Per-run one-element float32 output scale. Defaults to the plan.

  • validate (bool) – Validate the complete runtime contract. Disable only when the caller guarantees dtype, device, shape, stride, alignment, values, aliasing, and lifetime. Defaults to True.

class flashinfer.attention.prims_ts.BatchPrefillPagedTSWrapper(kv_layout: Literal['HND'] = 'HND')

Compile and reuse packed-Q context attention over HND paged K/V.

plan accepts only static compilation geometry. All request metadata is supplied to run: cumulative Q offsets, a fixed row-strided K/V page table, and logical K/V lengths. Metadata values may change between launches while staying within the planned capacities. The wrapper owns no context workspace and does not transform the fixed page table.

Plan-time scalar defaults are stored as one-element device tensors. run may replace either scale tensor without changing the compiled specialization. Plans default to a conservative dynamic-length contract; callers may instead promise exact-uniform packed lengths or no causal Q offset to compile one narrower specialization. Callers whose cache preparation zeroes every unused row in an active final V page may also compile out the defensive consumer-side V-tail clear. Validation reads request metadata back to the host, checks the length promises, and may synchronize; it does not inspect V-cache values. With caller-owned output, validate=False performs no allocation, metadata readback, or synchronization and is suitable for CUDA graph capture only when the caller already enforces the selected contract. Replanning invalidates captured graphs; prior launches and replays must finish before plan is called again. Keep the wrapper and all captured runtime tensors alive until every graph using the current plan is destroyed. Before running on a CUDA stream that is not already ordered after the planning stream, the caller must establish that dependency.

__init__(kv_layout: Literal['HND'] = 'HND') → None

Create an unplanned paged-context wrapper.

Parameters:

kv_layout ({"HND"}) – K/V layout. Only "HND" is supported.

plan(*, device: int | str | device, batch_size: int, max_seq_len_q: int, max_kv_len: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, q_dtype: dtype, kv_dtype: dtype, out_dtype: dtype | None = None, page_size: int = 32, mask_type: Literal['dense', 'causal'] = 'dense', window_left: int = -1, sm_scale: float | None = None, output_scale: float = 1.0, uniform_packed_lengths: bool = False, has_q_offset: bool = True, paged_v_tail_is_zero: bool = False) → None

Compile one reusable specialization from explicit static geometry.

Sequence lengths and page indices are deliberately absent. By default, the plan supports dynamic positive lengths bounded by max_seq_len_q and max_kv_len. uniform_packed_lengths=True instead promises that every runtime Q length equals max_seq_len_q and every runtime K/V length equals max_kv_len. has_q_offset=False promises that every causal request has Sq == Sk. Dense attention ignores and canonicalizes the latter flag. paged_v_tail_is_zero=True separately promises that the unused rows following each request’s logical K/V length in its active final V page are zero. The selected contract compiles exactly one specialization. run(validate=True) checks the length promises but cannot cheaply inspect V-cache contents; callers always own the V-tail promise. Violating a selected promise can produce incorrect results or invalid memory accesses.

batch_size is exact. Runtime page-table rows must expose at least ceil(max_kv_len / page_size) columns, including inactive padding columns. Calling plan again replaces plan-owned tensors and invalidates CUDA graphs captured from the previous plan. Complete every prior launch and replay before replanning. Planning allocates and compiles; complete it before CUDA Graph capture.

Parameters:
  • device (int, str, or torch.device) – CUDA device on which the specialization is compiled and run.

  • batch_size (int) – Exact number of requests in every run.

  • max_seq_len_q (int) – Maximum query length of any request.

  • max_kv_len (int) – Maximum logical K/V length of any request.

  • num_qo_heads (int) – Number of query/output heads.

  • num_kv_heads (int) – Number of key/value heads. num_qo_heads must be divisible by this value.

  • head_dim (int) – Query, key, value, and output head dimension.

  • q_dtype (torch.dtype) – Query dtype.

  • kv_dtype (torch.dtype) – Key/value cache dtype. It must currently equal q_dtype.

  • out_dtype (torch.dtype, optional) – Output dtype; defaults to q_dtype.

  • page_size (int) – Number of K/V tokens stored in each page.

  • mask_type ({"dense", "causal"}) – Attention mask mode.

  • window_left (int) – Left sliding-window extent, or -1 to disable the window.

  • sm_scale (float, optional) – Softmax scale; defaults to the inverse square root of head size.

  • output_scale (float) – Scale applied to the attention output.

  • uniform_packed_lengths (bool) – Whether every run has Q lengths exactly max_seq_len_q and K/V lengths exactly max_kv_len. Defaults to False.

  • has_q_offset (bool) – Whether causal runs may have a nonzero bottom-right Q offset Sk - Sq. Setting this to False promises Sq == Sk for every request. Dense attention ignores this flag. Defaults to True.

  • paged_v_tail_is_zero (bool) – Whether unused rows in every active final V-cache page are guaranteed to contain zero. Setting this to True compiles out the consumer-side V-tail clear. Defaults to False.

run(q: Tensor, k_cache: Tensor, v_cache: Tensor, qo_indptr: Tensor, block_tables: Tensor, seq_lens_kv: Tensor, *, out: Tensor | None = None, scale_softmax_log2: Tensor | None = None, output_scale: Tensor | None = None, validate: bool = True) → Tensor

Launch the compiled specialization with per-run request metadata.

Metadata must use int32 CUDA storage. qo_indptr is compact with shape [B + 1]; block_tables has shape [B, C], unit column stride, and row stride at least C; and seq_lens_kv is compact with shape [B]. Their values must satisfy the capacities documented by plan:

  • qo_indptr starts at zero, increases strictly, ends at the packed Q extent, and has deltas no greater than max_seq_len_q;

  • C is at least ceil(max_kv_len / page_size);

  • seq_lens_kv entries are positive and do not exceed max_kv_len;

  • causal launches additionally satisfy Sq[b] <= Sk[b]; and

  • every active page-table entry is a valid physical page ID in both separate, isomorphic K and V pools. Padding entries beyond the logical K/V length are not dereferenced.

When plan selected uniform_packed_lengths=True, every Q delta must equal max_seq_len_q and every K/V length must equal max_kv_len. For a causal plan with has_q_offset=False, every request must additionally satisfy Sq[b] == Sk[b].

With validate=True, the host reads these values and may synchronize. Metadata may change only between completed launches or graph replays; CUDA graph capture additionally requires stable tensor shapes, strides, and addresses plus validate=False. That mode does not verify the plan-time metadata promises; the caller is responsible for preserving them across every replay.

Parameters:
  • q (torch.Tensor) – Runtime packed query tensor matching the plan.

  • k_cache (torch.Tensor) – Runtime HND key page pool matching the plan.

  • v_cache (torch.Tensor) – Runtime HND value page pool matching the plan.

  • qo_indptr (torch.Tensor) – Cumulative packed-query offsets with shape [B + 1].

  • block_tables (torch.Tensor) – Physical page IDs with shape [B, C]. Column stride must be one and row stride may exceed C.

  • seq_lens_kv (torch.Tensor) – Logical K/V lengths with shape [B].

  • out (torch.Tensor, optional) – Caller-owned output tensor. A new tensor is allocated when omitted.

  • scale_softmax_log2 (torch.Tensor, optional) – Per-run one-element float32 softmax scale in base-2 form. Defaults to the scale retained by the plan.

  • output_scale (torch.Tensor, optional) – Per-run one-element float32 output scale. Defaults to the scale retained by the plan.

  • validate (bool) – Run explicit storage, shape, dtype, device, scale, output, and aliasing validators. Disable only when the caller guarantees the complete runtime contract. Defaults to True.

Returns:

The packed attention output.

Return type:

torch.Tensor

FMHA Decode

batch_decode_with_paged_kv_cache(q, ...[, ...])

One-shot fixed or packed-Q paged decode from fixed page tables.

get_prims_ts_batch_decode_workspace_size(...)

Return caller-workspace bytes for one automatic FMHA policy.

prims_ts_batch_decode_with_kv_cache(query, ...)

Launch fixed or packed-Q native page-table decode with caller scratch.

class flashinfer.attention.prims_ts.BatchDecodePagedTSWrapper(kv_layout: Literal['HND'] = 'HND')

Plan static paged-decode capacity and bind request metadata.

A plan fixes device, batch size, geometry, dtypes, query storage mode, and maximum K/V capacity. Sequence lengths are supplied either once to plan(), which retains an owned CUDA copy, or to every run(). Every run supplies a row-strided page table; packed-query plans also consume per-run Q offsets.

One wrapper supports one ordered execution lane. Concurrent streams or graph replays require separate wrappers and workspace buffers.

__init__(kv_layout: Literal['HND'] = 'HND') → None

Initialize an unplanned wrapper with one static K/V layout.

Parameters:

kv_layout ({"HND"}) – Layout of the paged K/V cache. Only "HND" is supported.

plan(device: int | str | device, batch_size: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, max_kv_len: int, *, max_seq_len_q: int = 1, packed_query: bool = False, q_data_type: dtype = torch.float16, kv_data_type: dtype | None = None, o_data_type: dtype | None = None, mask_type: Literal['dense', 'causal'] = 'dense', window_left: int = -1, seq_lens: Sequence[int] | Tensor | None = None, workspace_buffer: Tensor | None = None) → None

Compile one static-capacity plan and optionally own sequence lengths.

max_seq_len_q is the exact fixed Q length when packed_query is false and the per-request capacity when packed_query is true. A packed run supplies qo_indptr and compact [total_q, Hq, D] storage.

seq_lens optionally fixes the per-request lengths for the lifetime of this plan. Planning validates the host values, retains them, and creates an owned CUDA copy used by every launch. In this mode, run() requires its seq_lens argument to be None. When omitted, both K/V prefix and length handling compile dynamically and every run must supply a CUDA length tensor.

workspace_buffer is caller-owned scratch for this plan. It is allocated when omitted, initialized during planning, retained by the frozen plan state, and never reset by run.

Parameters:
  • device (int, str, or torch.device) – CUDA device on which the specialization is compiled and run.

  • batch_size (int) – Exact number of requests in every run.

  • num_qo_heads (int) – Number of query/output heads.

  • num_kv_heads (int) – Number of K/V heads.

  • head_dim (int) – Query, key, value, and output head dimension.

  • page_size (int) – Number of K/V tokens stored in each page.

  • max_kv_len (int) – Per-request K/V length capacity used for policy selection, compilation, and workspace sizing.

  • max_seq_len_q (int) – Exact fixed Q length, or per-request capacity for packed Q. Defaults to 1.

  • packed_query (bool) – Select compact [total_q, Hq, D] query/output storage instead of fixed [B, SQ, Hq, D] storage. Fixed SQ1 storage is [B, Hq, D]. Defaults to False.

  • q_data_type (torch.dtype) – Query dtype used to compile the plan. Defaults to torch.float16.

  • kv_data_type (torch.dtype, optional) – K/V dtype used to compile the plan. Defaults to q_data_type.

  • o_data_type (torch.dtype, optional) – Output dtype used to compile the plan. Defaults to q_data_type.

  • mask_type ({"dense", "causal"}) – Attention mask mode. Defaults to "dense".

  • window_left (int) – Left sliding-window extent, or -1 to disable the window. A non-negative value requires causal masking.

  • seq_lens (Sequence[int] or torch.Tensor, optional) – Host-only per-request K/V lengths owned by the resulting plan and used for specialization. A tensor must be a one-dimensional CPU int32 or int64 tensor. The sequence must contain exactly batch_size positive values no larger than max_kv_len.

  • workspace_buffer (torch.Tensor, optional) – Caller-owned contiguous int8 or uint8 scratch on device. It must be 32-byte aligned and large enough for the selected plan. When omitted, planning allocates the buffer. The retained buffer is exclusive to one in-flight launch or graph replay.

run(q: Tensor, paged_kv_cache: Tensor | tuple[Tensor, Tensor], seq_lens: Tensor | None, block_tables: Tensor, *, qo_indptr: Tensor | None = None, bmm1_scale: float | None = None, bmm2_scale: float = 1.0, out: Tensor | None = None, validate: bool = True) → Tensor

Launch the current plan with one sequence-length owner.

When plan() received sequence lengths, seq_lens must be None and the plan’s owned CUDA copy is used. Otherwise, seq_lens must be supplied on every run. This ownership check is unconditional.

validate=True performs structural, value, and alias validation. It reads per-run metadata values back to the host. This is the safe public default. validate=False treats every run argument as a trusted binding, performs no explicit wrapper validation, and remains free of metadata device-to-host synchronization. K/V view selection, scale forwarding, and optional output allocation are unavoidable in both modes.

Packed plans require qo_indptr with B + 1 int32 offsets. Fixed plans require qo_indptr to be omitted. Per-run metadata tensors may change identity between ordered runs.

Parameters:
  • q (torch.Tensor) – Runtime fixed or packed query tensor matching the plan.

  • paged_kv_cache (torch.Tensor or tuple[torch.Tensor, torch.Tensor]) – Runtime combined or separate paged K/V storage matching the plan.

  • seq_lens (torch.Tensor, optional) – Per-run contiguous int32 CUDA K/V lengths with shape [B]. Required when omitted from plan() and otherwise required to be None.

  • block_tables (torch.Tensor) – Per-run int32 CUDA physical page IDs with shape [B, C]. Entries must be contiguous within each row; the row stride may be any value at least C. Inactive tail entries are ignored.

  • qo_indptr (torch.Tensor, optional) – Per-run cumulative query offsets with shape [B + 1]. Required for a packed-query plan and rejected for a fixed-query plan.

  • bmm1_scale (float, optional) – QK scaling factor. Defaults to the inverse square root of head_dim.

  • bmm2_scale (float) – Value/output scaling factor. Defaults to 1.0.

  • out (torch.Tensor, optional) – Caller-owned output tensor. A new tensor is allocated when omitted.

  • validate (bool) – Run explicit structural, value, and alias validation. Disable only when the caller guarantees the complete runtime contract. Sequence length ownership is enforced in either mode. Defaults to True.

Returns:

The fixed or packed attention output.

Return type:

torch.Tensor

Block-Sparse FMHA

block_sparse_attention(q, k, v, ...[, ...])

Plan and run one compact-BSHD block-sparse attention launch.

block_sparse_attention_with_paged_kv_cache(q, ...)

Plan and run one fixed-Q paged block-sparse attention launch.

class flashinfer.attention.prims_ts.BlockSparseTSWrapper

Plan and reuse compact-BSHD block-sparse attention launches.

Q is [B, Sq, Hq, D] and K/V are [B, Skv, Hkv, D]. Sparse rows are owned per batch, KV head, and query block, so every Q head in one grouped KV head consumes the same sparse row. A plan fixes geometry and a per-row capacity; every run supplies either BSR or a packed exact-block bitmask. Proxy-enabled plans additionally consume caller-owned K/V summaries, while an optional token mask applies only to exact routes. Callers must keep those tensors alive and immutable until the queued run or captured graph finishes using them. CUDA Graph capture pins plan-owned state only, so captured routing storage remains the caller’s responsibility.

One plan revision owns one mutable route workspace. Its runs must be ordered on one stream or externally synchronized; unordered concurrent runs require distinct wrappers.

__init__() → None
plan(batch_size: int, seq_len_q: int, seq_len_kv: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, q_block_size: int, kv_block_size: int, *, device: device | str | int, max_blocks_per_row: int, use_kv_valid_bits: bool, sparse_format: Literal['bsr', 'bitmask'] = 'bsr', use_proxy_routes: bool = False, mask_type: Literal['dense', 'causal'] = 'dense', q_data_type: dtype = torch.float16, kv_data_type: dtype | None = None, o_data_type: dtype | None = None) → None

Choose a legal profile and allocate reusable routing capacity.

The plan owns immutable geometry and a uniform route workspace, not a sparse pattern. max_blocks_per_row bounds each runtime sparse row in semantic kv_block_size blocks. sparse_format="bsr" consumes canonical CSR-style rows, while "bitmask" consumes packed exact- block bits. Enabling proxy routes represents unselected blocks through caller-provided K/V summaries and currently requires mask_type="dense". Exact-only plans continue to support causal masking. use_kv_valid_bits selects whether every run() must supply the shared batch token mask. Callers may pass different routing tensor identities and index extents to each run as long as they fit this declared capacity.

MHA, GQA, and MQA are supported with Hq / Hkv a power of two no greater than 32 and D=128. Q, K, V, and O use one matching torch.float16 or torch.bfloat16 dtype. Runtime tensor shapes are documented by run().

q_block_size may be any positive signed-Int32 value for which a physical Q tile stays within one BSR row. Equivalently, q_block_size * (Hq / Hkv) must be divisible by 8; therefore Q block sizes 1, 2, and 4 require GQA ratios of at least 8, 4, and 2, respectively. kv_block_size may be 8, 16, 32, or a positive multiple of 64. The Q tile groups complete Q-head groups and as many Q tokens as fit without crossing a semantic Q-block row, up to Q128; fine KV blocks cap this at a SWAPAB Q32 tile. Proxy routes reuse the same Q-tile, KV-route, and MMA geometry as exact routes, but currently use the direct scheduler because reusable planning cannot observe live exact-route work. Every run prepares its selected BSR or bitmask into compact, profile-selected fixed-width route metadata, and the attention core consumes only that metadata. This remains true when every KV block is selected; callers that know a pattern is dense should choose the dense FMHA API explicitly.

Planning does not inspect routing values and does not synchronize the host. Reusable runs trust those values; assertion-enabled CuTe DSL builds diagnose contract violations on device. The one-shot block_sparse_attention() entry point retains synchronous canonical-BSR inspection so it can derive the smallest semantic row bound for that call.

Concurrent plans are serialized; run keeps using the published state. One revision has one mutable route workspace, so its runs must be ordered on one stream or externally synchronized. Unordered concurrent runs require distinct wrappers.

run(q: Tensor, k: Tensor, v: Tensor, block_indptr: Tensor | None = None, block_indices: Tensor | None = None, *, exact_block_bits: Tensor | None = None, k_summary: Tensor | None = None, v_summary: Tensor | None = None, kv_valid_bits: Tensor | None = None, sm_scale: float | None = None, out: Tensor | None = None) → Tensor

Launch the current plan on the caller’s current CUDA stream.

q and out use compact [B, Sq, Hq, D] while k and v use compact [B, Skv, Hkv, D], with shapes and dtypes fixed by plan(). If supplied, out must match Q’s shape and the planned output dtype. The returned tensor is exactly out when one was supplied; otherwise it is a newly allocated compact BSHD tensor. Only O is returned; this PrimTS API does not return LSE. The launch is enqueued asynchronously on the caller’s current CUDA stream.

A BSR plan consumes compact Int32 block_indptr with shape [B, Hkv, ceil(Sq / q_block_size) + 1] and compact Int32 block_indices. A bitmask plan instead requires both BSR arguments to be None and consumes packed UInt32 exact_block_bits with shape [B, Hkv, ceil(Sq / q_block_size), ceil(num_kv_blocks / 32)]. Bit r of word w selects block 32 * w + r; final-word padding bits are ignored. A proxy plan additionally consumes compact k_summary and v_summary with shape [B, num_kv_blocks, Hkv, D]. K summaries are block means and V summaries are block sums; the final partial block covers only its structural tokens.

Every row must fit the planned semantic-block capacity. Reusable runs trust routing values. CuTe DSL assertions can diagnose violations when enabled before compilation; otherwise invalid values have undefined behavior and may access out of bounds. A masked plan requires kv_valid_bits with shape [B, ceil(Skv / 32)] and dtype UInt32; an unmasked plan requires None. The mask applies only to raw exact routes; proxy summaries and their represented-token mass remain caller-defined. Routing tensors may have different identities on every run.

Keep this wrapper alive until every captured CUDA Graph is destroyed.

Parameters:
  • q (torch.Tensor) – Compact query tensor [B, Sq, Hq, D] matching the plan.

  • k (torch.Tensor) – Compact key tensor [B, Skv, Hkv, D] matching the plan.

  • v (torch.Tensor) – Compact value tensor with the same shape, dtype, and strides as k.

  • block_indptr (torch.Tensor, optional) – Contiguous Int32 BSR row offsets with shape [B, Hkv, ceil(Sq / q_block_size) + 1]. Required by BSR plans.

  • block_indices (torch.Tensor, optional) – Contiguous Int32 semantic KV-block IDs referenced by block_indptr. Required by BSR plans.

  • exact_block_bits (torch.Tensor, optional) – Compact packed UInt32 exact-block bitmap required by bitmask plans.

  • k_summary (torch.Tensor, optional) – Per-block mean K tensor required by proxy plans.

  • v_summary (torch.Tensor, optional) – Per-block summed V tensor required by proxy plans.

  • kv_valid_bits (torch.Tensor, optional) – Contiguous UInt32 token-validity bitmap [B, ceil(Skv / 32)]. Supply it exactly when the plan enabled token validity bits.

  • sm_scale (float, optional) – Softmax scale. Defaults to 1 / sqrt(D).

  • out (torch.Tensor, optional) – Caller-owned compact output buffer [B, Sq, Hq, D] with the planned output dtype.

Returns:

The compact output tensor; identical to out when provided.

Return type:

torch.Tensor

class flashinfer.attention.prims_ts.BlockSparsePagedTSWrapper

Plan fixed capacity and run with entirely live request metadata.

__init__() → None
plan(batch_size: int, seq_len_q: int, max_seq_len_kv: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, q_block_size: int, kv_block_size: int, page_size: int, *, device: device | str | int, max_blocks_per_row: int, use_kv_valid_bits: bool, mask_type: Literal['dense', 'causal'] = 'dense', q_data_type: dtype = torch.float16, kv_data_type: dtype | None = None, o_data_type: dtype | None = None) → None

Plan fixed-Q geometry and a maximum variable-K capacity.

The plan stores no request metadata. Every run supplies live page-table offsets, page IDs, sequence lengths, sparse routes, and optional token bits. max_seq_len_kv fixes compilation and mask capacity. Attention consumes caller-owned live lengths directly, without a plan-owned copy or device-to-host validation.

q_block_size may be any positive signed-Int32 value satisfying q_block_size * (Hq / Hkv) % 8 == 0. This row-purity condition keeps every physical Q tile within exactly one logical BSR row. kv_block_size remains restricted to 8, 16, 32, or a positive multiple of 64.

run(q: Tensor, paged_kv_cache: Tensor | tuple[Tensor, Tensor], paged_kv_indptr: Tensor, paged_kv_indices: Tensor, seq_lens_kv: Tensor, block_indptr: Tensor, block_indices: Tensor, *, kv_valid_bits: Tensor | None = None, sm_scale: float | None = None, out: Tensor | None = None) → Tensor

Launch with live lengths, page tables, and sparse routes.

Q and O are compact [B, Sq, Hq, D] tensors, including Sq=1. The cache is either combined [P, 2, Hkv, page, D] or a (K, V) tuple whose members are [P, Hkv, page, D] with compact inner HND strides and arbitrary non-overlapping outer page strides.

paged_kv_indptr is compact Int32 [B + 1]; paged_kv_indices is compact Int32 with capacity at least its live final offset; and seq_lens_kv is compact Int32 [B]. All values are read on device. The caller must keep every dense length in [1, max_seq_len_kv] and every causal length in [Sq, max_seq_len_kv]. paged_kv_indptr must start at zero and contain bounded, monotone rows with at least ceil(seq_lens_kv[b] / page_size) entries. Every page ID in its live prefix must lie in [0, P). Each BSR row must contain strictly increasing, unique block IDs whose final block starts before that request’s live K/V length, and its width must not exceed the planned max_blocks_per_row. Reusable runs trust all of these device-side values; assertion-enabled CuTe DSL builds diagnose violations encountered while preparing selected routes, while default builds leave invalid values undefined and may access out of bounds. No runtime value validation synchronizes or copies metadata to the host.

In eager execution, record_stream extends the allocator lifetime of every launch tensor (Q, normalized K/V, O, BSR metadata, token bits, and page metadata) through the asynchronous run. It does not make concurrent mutation safe. During CUDA Graph capture and replay, the caller must keep Q, the cache, O, and all runtime metadata alive and unmodified until replay completes; do not release or overwrite them while a replay is outstanding. The wrapper and its captured plan state must also outlive the graph. One plan revision owns mutable route scratch; unordered concurrent runs require separate wrapper instances.

Parameters:
  • q (torch.Tensor) – Compact query tensor [B, Sq, Hq, D] matching the plan.

  • paged_kv_cache (PagedKVCache) – Either a combined cache [P, 2, Hkv, page_size, D] or a (K, V) tuple whose tensors are [P, Hkv, page_size, D].

  • paged_kv_indptr (torch.Tensor) – Contiguous Int32 live request offsets with shape [B + 1].

  • paged_kv_indices (torch.Tensor) – Contiguous Int32 physical-page ID capacity.

  • seq_lens_kv (torch.Tensor) – Contiguous Int32 live logical K/V lengths with shape [B]. Values must satisfy the dense or causal bounds above.

  • block_indptr (torch.Tensor) – Contiguous Int32 BSR row offsets with shape [B, Hkv, ceil(Sq / q_block_size) + 1].

  • block_indices (torch.Tensor) – Contiguous Int32 logical KV-block IDs referenced by block_indptr.

  • kv_valid_bits (torch.Tensor, optional) – Contiguous UInt32 logical-token validity bitmap [B, ceil(max_seq_len_kv / 32)]. Supply it exactly when the plan enabled token validity bits.

  • sm_scale (float, optional) – Softmax scale. Defaults to 1 / sqrt(D).

  • out (torch.Tensor, optional) – Caller-owned compact output buffer [B, Sq, Hq, D] with the planned output dtype.

Returns:

The compact output tensor; identical to out when provided.

Return type:

torch.Tensor

MLA Decode

batch_mla_decode_with_paged_kv_cache(query, ...)

One-shot convenience wrapper for fixed or packed-query MLA decode.

get_prims_ts_batch_mla_decode_workspace_size(...)

Return caller-workspace bytes for one automatic MLA policy.

prims_ts_batch_mla_decode_with_kv_cache(...)

Launch fixed or packed-query paged MLA decode with caller-owned scratch.

class flashinfer.attention.prims_ts.BatchMLADecodePagedTSWrapper

Compile and reuse task-scheduled paged MLA decode launches.

__init__() → None

Initialize an unplanned task-scheduled paged-MLA wrapper.

plan(device: int | str | device, batch_size: int, num_heads: int, kv_lora_rank: int, qk_rope_head_dim: int, page_size: int, max_kv_len: int, *, max_seq_len_q: int, packed_query: bool, q_data_type: dtype, kv_data_type: dtype, o_data_type: dtype, mask_type: Literal['dense', 'causal'] = 'causal', workspace_buffer: Tensor | None = None) → None

Compile one static MLA shape and bind its reusable workspace.

Planning consumes only compile-time geometry, capacity, dtype, and storage-mode inputs. Request metadata belongs exclusively to run() and is never retained by the wrapper. A successful re-plan atomically replaces the previous immutable state; a failed re-plan leaves that state usable.

packed_query=False selects fixed [B, SQ, H, 576] query storage, where SQ is exactly max_seq_len_q. packed_query=True selects [total_q, H, 576] storage with per-run cumulative offsets supplied to every run. max_seq_len_q is then the per-request capacity. Individual packed requests may be empty; an all-empty run returns its empty output without dispatching a GPU kernel.

If workspace_buffer is omitted, the plan allocates private scratch. A workspace is mutable and exclusive to one in-flight launch or graph replay. Warm the plan before graph capture, and call run with validate=False inside compiled or captured regions.

Parameters:
  • device (int, str, or torch.device) – CUDA device on which the plan will execute.

  • batch_size (int) – Exact runtime request count.

  • num_heads (int) – Static MLA head geometry and K/V page size.

  • kv_lora_rank (int) – Static MLA head geometry and K/V page size.

  • qk_rope_head_dim (int) – Static MLA head geometry and K/V page size.

  • page_size (int) – Static MLA head geometry and K/V page size.

  • max_kv_len (int) – Static per-request K/V and Q capacities.

  • max_seq_len_q (int) – Static per-request K/V and Q capacities.

  • packed_query (bool) – Select packed rather than fixed query storage.

  • q_data_type (torch.dtype) – Query, K/V, and output dtypes used to compile the plan.

  • kv_data_type (torch.dtype) – Query, K/V, and output dtypes used to compile the plan.

  • o_data_type (torch.dtype) – Query, K/V, and output dtypes used to compile the plan.

  • mask_type ({"dense", "causal"}) – Attention mask mode.

  • workspace_buffer (torch.Tensor, optional) – Caller-owned contiguous int8 or uint8 scratch on device. It must be 32-byte aligned and large enough for the selected plan. When omitted, planning allocates the buffer. The retained buffer is exclusive to one in-flight launch or graph replay.

run(query: Tensor, kv_cache: Tensor, block_tables: Tensor, seq_lens: Tensor, *, qo_indptr: Tensor | None = None, bmm1_scale: float = 1.0, bmm2_scale: float = 1.0, out: Tensor | None = None, validate: bool = True) → Tensor

Launch the most recently planned MLA decode on the current stream.

block_tables and seq_lens are required per-run bindings. qo_indptr is required by a packed-query plan and rejected by a fixed-query plan. With validation enabled, tensor structure, metadata values, scales, aliases, and every static capacity are checked before launch. These checks synchronize metadata to the host. Set validate=False only after validating representative inputs, and use it for torch.compile or CUDA graph capture.

Parameters:
  • query (torch.Tensor) – Runtime fixed or packed query tensor matching the plan.

  • kv_cache (torch.Tensor) – Runtime compact paged latent K/V cache.

  • block_tables (torch.Tensor) – Runtime physical-page table with one inner-contiguous, non-overlapping row per request. Inter-row padding is accepted.

  • seq_lens (torch.Tensor) – Runtime K/V lengths with one element per request.

  • qo_indptr (torch.Tensor, optional) – Runtime cumulative query offsets for a packed-query plan.

  • bmm1_scale (float) – QK and value/output scaling factors.

  • bmm2_scale (float) – QK and value/output scaling factors.

  • out (torch.Tensor, optional) – Caller-owned output tensor. A new tensor is allocated when omitted.

  • validate (bool) – Enable explicit runtime validation. Defaults to True.

Returns:

The fixed or packed MLA attention output.

Return type:

torch.Tensor

flashinfer.decode

Single Request Decoding

single_decode_with_kv_cache()

Decode attention with KV Cache for single request, return attention output.

single_decode_with_kv_cache_with_jit_module(...)

Single-request decode using a pre-compiled JIT module.

Batch Decoding

cudnn_batch_decode_with_kv_cache(q, k_cache, ...)

Batched decode attention with paged KV cache, backed by cuDNN SDPA.

trtllm_batch_decode_with_kv_cache(query, ...)

xqa_batch_decode_with_kv_cache(query, ...[, ...])

DCP Speculative Decode Workspace

The native Cake FMHA DCP speculative route of flashinfer.decode.trtllm_batch_decode_with_kv_cache() uses caller-owned scratch buffers so a prewarmed invocation can be captured in a CUDA Graph. The production D256 FP8/page64 ratio-16 profile supports speculative query lengths 1 through 8 and passes head_dim=256 to the workspace-size helper; D128 remains the default.

get_dcp_spec_workspace_size_bytes(...[, ...])

Bytes for Cake FMHA Split-KV BF16 partial-O and FP32 partial-LSE scratch.

get_dcp_spec_counter_bytes(batch_size, ...)

Bytes for the v4 completion tickets, zeroed once then self-reset.

class flashinfer.decode.BatchDecodeWithPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, use_tensor_cores: bool = False, paged_kv_indptr_buffer: Tensor | None = None, paged_kv_indices_buffer: Tensor | None = None, paged_kv_last_page_len_buffer: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None)

Wrapper class for decode attention with paged kv-cache (first proposed in vLLM) for batch of requests.

Check our tutorial for page table layout.

Examples

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 8
>>> head_dim = 128
>>> max_num_pages = 128
>>> page_size = 16
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> decode_wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> kv_page_indices = torch.arange(max_num_pages).int().to("cuda:0")
>>> kv_page_indptr = torch.tensor(
...     [0, 17, 29, 44, 48, 66, 100, 128], dtype=torch.int32, device="cuda:0"
... )
>>> # 1 <= kv_last_page_len <= page_size
>>> kv_last_page_len = torch.tensor(
...     [1, 7, 14, 4, 3, 1, 16], dtype=torch.int32, device="cuda:0"
... )
>>> kv_cache_at_layer = [
...     torch.randn(
...         max_num_pages, 2, page_size, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> # create auxiliary data structures for batch decode attention
>>> decode_wrapper.plan(
...     kv_page_indptr,
...     kv_page_indices,
...     kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     pos_encoding_mode="NONE",
...     data_type=torch.float16
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = torch.randn(batch_size, num_qo_heads, head_dim).half().to("cuda:0")
...     kv_cache = kv_cache_at_layer[i]
...     # compute batch decode attention, reuse auxiliary data structures for all layers
...     o = decode_wrapper.run(q, kv_cache)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([7, 64, 128])

Note

To accelerate computation, FlashInfer’s batch decode attention creates some auxiliary data structures, these data structures can be reused across multiple batch decode attention calls (e.g. different Transformer layers). This wrapper class manages the lifecycle of these data structures.

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, use_tensor_cores: bool = False, paged_kv_indptr_buffer: Tensor | None = None, paged_kv_indices_buffer: Tensor | None = None, paged_kv_last_page_len_buffer: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None) → None

Constructor of BatchDecodeWithPagedKVCacheWrapper.

Parameters:
  • float_workspace_buffer (torch.Tensor. Must be initialized to 0 for its first use.) – The user reserved float workspace buffer used to store intermediate attention results in the split-k algorithm. The recommended size is 128MB, the device of the workspace buffer should be the same as the device of the input tensors. The buffer must be 16-byte aligned; tensors created by torch.empty satisfy this on supported devices.

  • kv_layout (str) – The layout of the input k/v tensors, could be either NHD or HND.

  • use_cuda_graph (bool) – Whether to enable CUDAGraph for batch decode attention, if enabled, the auxiliary data structures will be stored as the provided buffers. The batch_size cannot change during the lifecycle of this wrapper when CUDAGraph is enabled.

  • use_tensor_cores (bool) – Whether to use tensor cores for the computation. Will be faster for large group size in grouped query attention. Defaults to False.

  • paged_kv_indptr_buffer (Optional[torch.Tensor]) – The user reserved buffer on GPU to store the indptr of the paged kv cache, the size of the buffer should be [batch_size + 1]. Only needed when use_cuda_graph is True.

  • paged_kv_indices_buffer (Optional[torch.Tensor]) – The user reserved buffer on GPU to store the page indices of the paged kv cache, should be large enough to store the maximum number of page indices (max_num_pages) during the lifecycle of this wrapper. Only needed when use_cuda_graph is True.

  • paged_kv_last_page_len_buffer (Optional[torch.Tensor]) – The user reserved buffer on GPU to store the number of entries in the last page, the size of the buffer should be [batch_size]. Only needed when use_cuda_graph is True.

  • backend (str) – The implementation backend, could be auto/fa2/fa3/trtllm-gen/ cute-dsl or prims-ts. Defaults to auto. If set to auto, the wrapper will automatically choose the backend based on the device architecture and kernel availability. The cute-dsl backend uses the CuTe DSL GQA decode kernel for Blackwell (SM100+) and only supports a subset of features (equal head_dim_qk/vo, no RoPE/ALiBi/soft-cap). The prims-ts backend uses the task-scheduled decode kernel on SM100a/SM103a. It is the only backend that accepts is_causal=False with q_len_per_req > 1. It requires kv_layout="HND" and does not support use_cuda_graph=True. Manual CUDA graph capture binds to one completed plan; recapture after re-planning. See the PrimTS decode guide for metadata requirements.

  • jit_args (Optional[List[Any]]) – If provided, the wrapper will use the provided arguments to create the JIT module, otherwise, the wrapper will use default attention implementation.

plan(indptr: Tensor, indices: Tensor, last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, *deprecated_positional_args: Any, **kwargs: Any) → None

Plan batch decode for given problem specification.

Parameters:
  • indptr (torch.Tensor) – The indptr of the paged kv cache, shape: [batch_size + 1], dtype: torch.int32

  • indices (torch.Tensor) – The page indices of the paged kv cache, shape: [indptr[-1]], dtype: torch.int32

  • last_page_len (torch.Tensor) – The number of entries in the last page of each request in the paged kv cache, shape: [batch_size], dtype: torch.int32

  • num_qo_heads (int) – The number of query/output heads

  • num_kv_heads (int) – The number of key/value heads

  • head_dim (int) – The dimension of the heads

  • page_size (int) – The page size of the paged kv cache

  • pos_encoding_mode (str) – The position encoding applied inside attention kernels, could be NONE/ROPE_LLAMA (LLAMA style rotary embedding) /ALIBI. Defaults to NONE.

  • window_left (int) – The left (inclusive) window size for the attention window, when set to -1, the window size will be set to the full length of the sequence. Defaults to -1.

  • window_right (int) – The right (inclusive) window size for the attention window. -1 disables the right bound. Defaults to 0. Currently window_right != 0 only supported by the cute-dsl backend.

  • logits_soft_cap (Optional[float]) – The attention logits soft capping value (used in Gemini, Grok and Gemma-2, etc.), if not provided, will be set to 0. If greater than 0, the logits will be capped according to formula: \(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\), where \(x\) is the input logits.

  • q_data_type (Optional[Union[str, torch.dtype]]) – The data type of the query tensor, defaults torch.float16.

  • kv_data_type (Optional[Union[str, torch.dtype]]) – The data type of the key/value tensor. If None, will be set to q_data_type. Defaults to None.

  • o_data_type (Optional[Union[str, torch.dtype]]) – The data type of the output tensor. If None, will be set to q_data_type. For FP8 inputs, this should typically be set to torch.float16 or torch.bfloat16.

  • data_type (Optional[Union[str, torch.dtype]]) – The data type of both the query and key/value tensors. Defaults to torch.float16. data_type is deprecated, please use q_data_type and kv_data_type instead.

  • sm_scale (Optional[float]) – Softmax scale. If None, defaults to 1 / sqrt(head_dim). Cached on the wrapper and reused at run() time.

  • rope_scale (Optional[float]) – Scale factor applied during RoPE interpolation. Only consulted when pos_encoding_mode != "NONE". Defaults to 1.0 when None.

  • rope_theta (Optional[float]) – Base value for the RoPE frequencies. Only consulted when pos_encoding_mode != "NONE". Defaults to 1e4 when None.

  • non_blocking (bool) – Whether to copy the input tensors to the device asynchronously, defaults to True.

  • seq_lens (Optional[torch.Tensor]) – A uint32 1D tensor indicating the K/V sequence length of each prompt, with shape [batch_size]. If omitted, lengths are derived from indptr, last_page_len, and page_size.

  • block_tables (Optional[torch.Tensor]) – A uint32 2D tensor indicating the physical page IDs of each prompt, with shape [batch_size, max_num_blocks_per_seq]. If omitted, backends that use block tables construct them from the CSR inputs during planning.

  • fixed_split_size (Optional[int],) – The fixed split size for FA2 split-kv decode, in pages. Only supported by tensor core decode for now. Recommend setting to the average sequence length of your workload. When enabled for FA2, will lead to deterministic softmax score reduction in the merge_states kernel, and therefore batch-size invariant outputs. See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ Note that compatibility with CUDA graph is NOT guaranteed, as even when bs is fixed, kv seq len can change and lead to a varied number of launched CTAs.

  • disable_split_kv (bool,) – Whether to disable the split-kv for determinism in CUDA Graph, defaults to False.

  • q_len_per_req (int) – The number of query tokens per request. Defaults to 1. q_len_per_req > 1 is currently supported on the fa2 tensor-core backend and natively by trtllm-gen, cute-dsl, and prims-ts. Call plan() again before changing this value. Under use_cuda_graph, this value is part of the frozen shape (like the batch size): once the wrapper has been planned, re-planning with a different value raises.

  • is_causal (Optional[bool]) – Whether the mask is causal within each request block. Defaults to None, which derives it from q_len_per_req > 1. Only the prims-ts backend honors a value that differs from that default; the other backends raise.

Note

The plan() method should be called before any run() or run_return_lse() calls, auxiliary data structures will be created during this call and cached for multiple run calls. After a failed plan() call, call plan() successfully again before calling run() or run_return_lse().

The num_qo_heads must be a multiple of num_kv_heads. If num_qo_heads is not equal to num_kv_heads, the function will use grouped query attention.

The plan() method cannot be used in Cuda Graph or in torch.compile.

Optional arguments after page_size are accepted positionally for backward compatibility, but that calling convention is deprecated and scheduled for removal in a future release. Pass them by keyword instead. window_right is keyword-only.

prewarm_paged_kv_stride_variant(variant: str = 'independent') → None

Load a lazy tensor-core paged-KV-stride variant after plan().

Call this method before CUDA graph capture when a planned standard FA2 tensor-core decode wrapper will receive K and V tensors with different data strides.

Parameters:

variant (str) – The paged-KV-stride variant to prewarm. The only supported value is "independent" (the default), for K and V with different data strides.

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) → None

Reset the workspace buffer.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The new float workspace buffer, the device of the new float workspace buffer should be the same as the device of the input tensors.

  • int_workspace_buffer (torch.Tensor) – The new int workspace buffer, the device of the new int workspace buffer should be the same as the device of the input tensors.

run(q: Tensor, paged_kv_cache: torch.Tensor | Tuple[torch.Tensor, torch.Tensor], *args, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None, window_left: int | None = None, sinks: Tensor | None = None, q_len_per_req: int | None = None, skip_softmax_threshold_scale_factor: float | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None) → Tensor
run(q: Tensor, paged_kv_cache: torch.Tensor | Tuple[torch.Tensor, torch.Tensor], *args, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None, window_left: int | None = None, sinks: Tensor | None = None, q_len_per_req: int | None = None, skip_softmax_threshold_scale_factor: float | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None) → Tuple[Tensor, Tensor]

Compute batch decode attention between query and paged kv cache.

Parameters:
  • q (torch.Tensor) – The query tensor, shape: [batch_size * q_len_per_req, num_qo_heads, head_dim] On the trtllm-gen and cute-dsl backends the q_len_per_req implied by q.shape[0] may differ from the planned value; the fa2/fa3 tensor-core path requires it to match plan().

  • paged_kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) –

    The paged KV-Cache stored as a tuple of tensors or a single tensor:

    • a tuple (k_cache, v_cache) of 4-D tensors, each with shape: [max_num_pages, page_size, num_kv_heads, head_dim] if kv_layout is NHD, and [max_num_pages, num_kv_heads, page_size, head_dim] if kv_layout is HND.

    • a single 5-D tensor with shape: [max_num_pages, 2, page_size, num_kv_heads, head_dim] if kv_layout is NHD, and [max_num_pages, 2, num_kv_heads, page_size, head_dim] if kv_layout is HND. Where paged_kv_cache[:, 0] is the key-cache and paged_kv_cache[:, 1] is the value-cache.

  • *args – Additional arguments for the custom kernel.

  • q_scale (Optional[float]) – The calibration scale of query for fp8 input, if not provided, will be set to 1.0.

  • k_scale (Optional[float]) – The calibration scale of key for fp8 or nvfp4 input, if not provided, will be set to 1.0.

  • v_scale (Optional[float]) – The calibration scale of value for fp8 or nvfp4 input, if not provided, will be set to 1.0.

  • out (Optional[torch.Tensor]) – The output tensor, if not provided, will be allocated internally. Must be zero-init for cute-dsl backend.

  • lse (Optional[torch.Tensor]) – The log-sum-exp of attention logits, if not provided, will be allocated internally.

  • return_lse (bool) – Whether to return the logsumexp of attention scores, defaults to False.

  • enable_pdl (bool) – Whether to enable Programmatic Dependent Launch (PDL). See https://docs.nvidia.com/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization Only supported for >= sm90, and currently only for FA2 and CUDA core decode.

  • window_left (Optional[int]) – Per-call sliding-window bound. Must either be None (default, inherit the value from plan()) or equal the value passed to plan() (the kernel asserts this). Passing -1 to plan() disables the sliding window for the entire batch.

  • sinks (Optional[torch.Tensor]) – Per-head attention sink logits, shape [num_qo_heads]. When provided, sinks[head_idx] is appended to each row of the softmax denominator (Streaming-LLM / Attention-Sinks). The dtype requirement is backend-specific and validated by the underlying kernel; pass None to disable.

  • q_len_per_req (Optional[int]) – DEPRECATED — pass to plan() instead. When provided here, emits a DeprecationWarning and is used to validate the run-time value inferred from q.size(0). Scheduled for removal in a future release.

  • skip_softmax_threshold_scale_factor (Optional[float] = None) – threshold scale factor for skipping softmax operations. Providing a value for this parameter enables skip-softmax sparsity as described in: https://arxiv.org/abs/2512.12087 If no value is provided, then standard attention is used. Setting the threshold to a higher value generally increases kernel performance at the cost of accuracy degradation. The actual threshold value equals the provided threshold_scale_factor divided by the context length.

  • kv_cache_sf (Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]]) –

    Per-block scale factors for NVFP4 KV cache. Accepts the same formats as paged_kv_cache:

    • a tuple (k_scales, v_scales) of 4-D tensors, each with shape: [num_pages, page_size, num_kv_heads, head_dim // 16] if kv_layout is NHD, and [num_pages, num_kv_heads, page_size, head_dim // 16] if kv_layout is HND.

    • a single 5-D tensor with shape: [num_pages, 2, page_size, num_kv_heads, head_dim // 16] if kv_layout is NHD, and [num_pages, 2, num_kv_heads, page_size, head_dim // 16] if kv_layout is HND, where dim 1 holds k (index 0) and v (index 1) scales.

    Both tensors have dtype torch.float8_e4m3fn.

    Currently, NVFP4 KV supports fa2 and trtllm-gen backend.

Returns:

If return_lse is False, the attention output, shape: [batch_size, num_qo_heads, head_dim]. If return_lse is True, a tuple of two tensors:

  • attention output, shape: [batch_size, num_qo_heads, head_dim]

  • logsumexp of attention scores, shape: [batch_size, num_qo_heads].

Return type:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

workspace_size(indptr: Tensor, indices: Tensor, last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, pos_encoding_mode: str = 'NONE', window_left: int = -1, logits_soft_cap: float | None = None, q_data_type: str | dtype | None = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, data_type: str | dtype | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, block_tables: Tensor | None = None, seq_lens: Tensor | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False, q_len_per_req: int = 1) → Tuple[int, int]

Return the caller-owned workspace size required by plan().

The returned tuple is (float_workspace_size, int_workspace_size) in bytes. The inputs follow plan(); host-side planning tensors such as indptr and last_page_len are copied to CPU in the same way as plan(). The wrapper’s float workspace buffer is only used as the device/stream selector for the underlying module query. This method does not allocate buffers and does not mutate cached plan state.

Parameters:
  • indptr (torch.Tensor) – The indptr of the paged kv cache, shape: [batch_size + 1], dtype: torch.int32

  • indices (torch.Tensor) – The page indices of the paged kv cache, shape: [indptr[-1]], dtype: torch.int32

  • last_page_len (torch.Tensor) – The number of entries in the last page of each request in the paged kv cache, shape: [batch_size], dtype: torch.int32

  • num_qo_heads (int) – The number of query/output heads.

  • num_kv_heads (int) – The number of key/value heads.

  • head_dim (int) – The dimension of the heads.

  • page_size (int) – The page size of the paged kv cache.

  • pos_encoding_mode (str) – The position encoding applied inside attention kernels, could be NONE/ROPE_LLAMA (LLAMA style rotary embedding) /ALIBI. Defaults to NONE.

  • window_left (int) – The left (inclusive) window size for the attention window, when set to -1, the window size will be set to the full length of the sequence. Defaults to -1.

  • logits_soft_cap (Optional[float]) – The attention logits soft capping value (used in Gemini, Grok and Gemma-2, etc.), if not provided, will be set to 0.

  • q_data_type (Optional[Union[str, torch.dtype]]) – The data type of the query tensor, defaults to torch.float16.

  • kv_data_type (Optional[Union[str, torch.dtype]]) – The data type of the key/value tensor. If None, will be set to q_data_type.

  • o_data_type (Optional[Union[str, torch.dtype]]) – The data type of the output tensor. If None, will be set to q_data_type.

  • data_type (Optional[Union[str, torch.dtype]]) – Deprecated alias — sets both q_data_type and kv_data_type when they are not provided explicitly.

  • sm_scale (Optional[float]) – Softmax scale. If None, defaults to 1 / sqrt(head_dim).

  • rope_scale (Optional[float]) – Scale factor applied during RoPE interpolation. Defaults to 1.0 when None.

  • rope_theta (Optional[float]) – Base value for the RoPE frequencies. Defaults to 1e4 when None.

  • block_tables (Optional[torch.Tensor]) – Unused by this method; accepted for signature compatibility with plan().

  • seq_lens (Optional[torch.Tensor]) – A uint32 1D tensor indicating the kv sequence length of each prompt, shape: [batch_size].

  • fixed_split_size (Optional[int]) – The fixed split size for FA2 split-kv decode, in pages.

  • disable_split_kv (bool) – Whether to disable the split-kv for determinism in CUDA Graph. Defaults to False.

  • q_len_per_req (int) – The number of query tokens per request. Defaults to 1.

Returns:

(float_workspace_size, int_workspace_size) in bytes.

Return type:

Tuple[int, int]

Example

>>> float_bytes, int_bytes = wrapper.workspace_size(...)
>>> wrapper.reset_workspace_buffer(
...     torch.empty(float_bytes, dtype=torch.uint8, device="cuda"),
...     torch.empty(int_bytes, dtype=torch.uint8, device="cuda"),
... )
>>> wrapper.plan(...)
class flashinfer.decode.CUDAGraphBatchDecodeWithPagedKVCacheWrapper(workspace_buffer: Tensor, indptr_buffer: Tensor, indices_buffer: Tensor, last_page_len_buffer: Tensor, kv_layout: str = 'NHD', use_tensor_cores: bool = False)

CUDAGraph-compatible Wrapper class for decode attention with paged kv-cache (first proposed in vLLM) for batch of requests.

Note that this wrapper may not be as efficient as BatchDecodeWithPagedKVCacheWrapper because we won’t dispatch to different kernels for different batch sizes/sequence lengths/etc to accommodate the CUDAGraph requirement.

Check our tutorial for page table layout.

Note

The plan() method could not be captured by CUDAGraph.

__init__(workspace_buffer: Tensor, indptr_buffer: Tensor, indices_buffer: Tensor, last_page_len_buffer: Tensor, kv_layout: str = 'NHD', use_tensor_cores: bool = False) → None

Constructor of BatchDecodeWithPagedKVCacheWrapper.

Parameters:
  • workspace_buffer (torch.Tensor) – The user reserved workspace buffer on GPU used to store auxiliary data structures, recommended size is 128MB, the device of the workspace buffer should be the same as the device of the input tensors.

  • indptr_buffer (torch.Tensor) – The user reserved buffer on GPU to store the indptr of the paged kv cache, should be large enough to store the indptr of maximum batch size ([max_batch_size + 1]) during the lifecycle of this wrapper.

  • indices_buffer (torch.Tensor) – The user reserved buffer on GPU to store the page indices of the paged kv cache, should be large enough to store the maximum number of page indices (max_num_pages) during the lifecycle of this wrapper.

  • last_page_len_buffer (torch.Tensor) – The user reserved buffer on GPU to store the number of entries in the last page, should be large enough to store the maximum batch size ([max_batch_size]) during the lifecycle of this wrapper.

  • use_tensor_cores (bool) – Whether to use tensor cores for the computation. Will be faster for large group size in grouped query attention. Defaults to False.

  • kv_layout (str) – The layout of the input k/v tensors, could be either NHD or HND.

XQA

xqa(q, k_cache, v_cache, page_table, ...[, ...])

Apply attention with paged KV cache using XQA kernel. :param q: Query tensor with shape [batch_size, beam_width, num_q_heads, head_dim] if not using speculative decoding, or [batch_size, beam_width, q_seq_len, num_q_heads, head_dim] if using speculative decoding. q_seq_len is the number of speculative decoding tokens. Data type should be torch.float16 or torch.bfloat16. Now only beam_width 1 is supported. :type q: torch.Tensor :param k_cache: Paged K cache tensor with shape [num_pages, page_size, num_kv_heads, head_dim] if kv_layout is NHD, or [num_pages, num_kv_heads, page_size, head_dim] if kv_layout is HND. Data type should match query tensor or be torch.float8_e4m3fn, in which case xqa will run fp8 calculation. Should be the same data type as v_cache. When using NVFP4 KV, the data type is torch.uint8, and the last dimension should be head_dim / 2. :type k_cache: torch.Tensor :param v_cache: Paged V cache tensor with shape [num_pages, page_size, num_kv_heads, head_dim] if kv_layout is NHD, or [num_pages, num_kv_heads, page_size, head_dim] if kv_layout is HND. Data type should match query tensor or be torch.float8_e4m3fn, in which case xqa will run fp8 calculation. Should be the same data type as k_cache. When using NVFP4 KV, the data type is torch.uint8, and the last dimension should be head_dim / 2. :type v_cache: torch.Tensor :param k_sf_cache: Optional scale factor cache tensor for the K cache. Use when NVFP4 KV is used. Expected shape is [num_pages, page_size, num_kv_heads, head_dim / 16] if kv_layout is NHD, or [num_pages, num_kv_heads, page_size, head_dim / 16] if kv_layout is HND. Should be the same data type as v_sf_cache. Data type should be torch.uint8. :type k_sf_cache: Optional[torch.Tensor] :param v_sf_cache: Optional scale factor cache tensor for the V cache. Use when NVFP4 KV is used. Expected shape is [num_pages, page_size, num_kv_heads, head_dim / 16] if kv_layout is NHD, or [num_pages, num_kv_heads, page_size, head_dim / 16] if kv_layout is HND. Should be the same data type as k_sf_cache. Data type should be torch.uint8. :type v_sf_cache: Optional[torch.Tensor] :param page_table: Page table tensor with shape batch_size, nb_pages_per_seq. Data type should be torch.int32. K and V share the same table. :type page_table: torch.Tensor :param seq_lens: Sequence lengths tensor with shape [batch_size, beam_width]. Data type should be torch.uint32. :type seq_lens: torch.Tensor :param output: Output tensor with shape that matches the query tensor. Data type should match query tensor or kv tensor. This tensor will be modified in-place. :type output: torch.Tensor :param workspace_buffer: Workspace buffer for temporary computations. Data type should be torch.uint8. :type workspace_buffer: torch.Tensor :param semaphores: Semaphore buffer for synchronization. Data type should be torch.uint32. :type semaphores: torch.Tensor :param num_kv_heads: Number of key-value heads in the attention mechanism. :type num_kv_heads: int :param page_size: Size of each page in the paged KV cache. Must be one of [16, 32, 64, 128]. :type page_size: int :param sinks: Attention sink values with shape [num_kv_heads, head_group_ratio]. Data type should be torch.float32. If None, no attention sinks are used. :type sinks: Optional[torch.Tensor], default=None :param q_scale: Scale factor for query tensor. :type q_scale: Union[float, torch.Tensor], default=1.0 :param kv_scale: Scale factor for KV cache. :type kv_scale: Union[float, torch.Tensor], default=1.0 :param sliding_win_size: Sliding window size for attention. If 0, no sliding window is used. With speculative decoding (q_seq_len > 1), the window is applied per draft row, assuming draft tokens occupy consecutive positions at the end of the sequence (linear chains, causal or full masks); tree-structured drafts are not supported with sliding window. :type sliding_win_size: int, default=0 :param kv_layout: The layout of the KV cache. Can be either NHD or HND. :type kv_layout: str, default="NHD" :param sm_count: Number of streaming multiprocessors to use. If None, will be inferred from the device. :type sm_count: Optional[int], default=None :param enable_pdl: Whether to enable PDL (Persistent Data Loader) optimization. If None, will be set to True if hardware supports it. :type enable_pdl: Optional[bool], default=None :param rcp_out_scale: Reciprocal of output scale factor. :type rcp_out_scale: float, default=1.0 :param q_seq_len: Query sequence length. When > 1, enables speculative decoding mode. With ragged Q (q_cu_seq_lens set), this is the maximum draft length across the batch. :type q_seq_len: int, default=1 :param mask: Draft-block attention mask for speculative decoding mode (when q_seq_len > 1). Shape: [batch_size, q_seq_len,              mask_size_per_row] where mask_size_per_row = ((q_seq_len + 31) //              32) * 2, or [total_q_tokens, mask_size_per_row] with ragged Q, where each request's rows cover its own draft length. Data type should be torch.uint16 (bit-packed format, aligned to 32 bits). :type mask: Optional[torch.Tensor], default=None :param q_cu_seq_lens: Cumulative draft lengths [batch_size + 1] (torch.int32 or torch.uint32, on device) enabling ragged Q: requests may have different draft lengths, including 0 for a request with no draft tokens this step. When set, q and output are packed as [total_q_tokens, num_q_heads, head_dim], q_seq_len must be the maximum draft length, and mask rows are packed by the same cumulative offsets. :type q_cu_seq_lens: Optional[torch.Tensor], default=None.

xqa_mla(q, k_cache, v_cache, page_table, ...)

Apply attention with paged KV cache using XQA MLA (Multi-Head Latent Attention) kernel. :param q: Query tensor with shape [batch_size, beam_width, num_q_heads, head_dim]. Data type should be torch.float8_e4m3fn. Now only beam_width 1 is supported. :type q: torch.Tensor :param k_cache: Paged K cache tensor with shape [total_num_cache_heads, head_dim]. Data type should be torch.float8_e4m3fn :type k_cache: torch.Tensor :param v_cache: Paged V cache tensor with shape [total_num_cache_heads, head_dim]. Data type should be torch.float8_e4m3fn :type v_cache: torch.Tensor :param page_table: Page table tensor with shape batch_size, nb_pages_per_seq. Data type should be torch.int32. K and V share the same table. :type page_table: torch.Tensor :param seq_lens: Sequence lengths tensor with shape [batch_size, beam_width]. Data type should be torch.uint32. :type seq_lens: torch.Tensor :param output: Output tensor with shape [batch_size, beam_width, num_q_heads, head_dim]. Data type should be torch.bfloat16. This tensor will be modified in-place. :type output: torch.Tensor :param workspace_buffer: Workspace buffer for temporary computations. Data type should be torch.uint8. :type workspace_buffer: torch.Tensor :param semaphores: Semaphore buffer for synchronization. Data type should be torch.uint32. :type semaphores: torch.Tensor :param page_size: Size of each page in the paged KV cache. Must be one of [16, 32, 64, 128]. :type page_size: int :param q_scale: Scale factor for query tensor. :type q_scale: Union[float, torch.Tensor], default=1.0 :param kv_scale: Scale factor for KV cache. :type kv_scale: Union[float, torch.Tensor], default=1.0 :param sm_count: Number of streaming multiprocessors to use. If None, will be inferred from the device. :type sm_count: Optional[int], default=None :param enable_pdl: Whether to enable PDL (Persistent Data Loader) optimization. If None, will be set to True if hardware supports it. :type enable_pdl: Optional[bool], default=None.

flashinfer.prefill

Attention kernels for prefill & append attention in both single request and batch serving setting.

Single Request Prefill/Append Attention

single_prefill_with_kv_cache()

Prefill/Append attention with KV cache for single request, return the attention output.

single_prefill_with_kv_cache_return_lse(q, k, v)

Convenience wrapper for single_prefill_with_kv_cache() that always returns LSE.

single_prefill_with_kv_cache_with_jit_module(...)

Single-request prefill / append attention using a pre-compiled JIT module.

Batch Prefill/Append Attention

cudnn_batch_prefill_with_kv_cache(q, ...[, ...])

Batched prefill attention with paged KV cache, backed by cuDNN SDPA.

trtllm_batch_context_with_kv_cache(query, ...)

trtllm_ragged_attention_deepseek(query, key, ...)

fmha_v2_prefill_deepseek(query, key, value, ...)

trtllm_fmha_v2_prefill(qkv, input_layout, ...)

TRT-LLM FMHAv2 prefill attention.

fmha_v2_prefill_sm120(query, key, value, ...)

Run SM120 FMHA v2 with contiguous, separate Q/K/V tensors.

class flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None, variant_owns_mask: bool = False)

Wrapper class for prefill/append attention with paged kv-cache for batch of requests.

Check our tutorial for page table layout.

Example

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 16
>>> head_dim = 128
>>> max_num_pages = 128
>>> page_size = 16
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> prefill_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> nnz_qo = 100
>>> qo_indptr = torch.tensor(
...     [0, 33, 44, 55, 66, 77, 88, nnz_qo], dtype=torch.int32, device="cuda:0"
... )
>>> paged_kv_indices = torch.arange(max_num_pages).int().to("cuda:0")
>>> paged_kv_indptr = torch.tensor(
...     [0, 17, 29, 44, 48, 66, 100, 128], dtype=torch.int32, device="cuda:0"
... )
>>> # 1 <= paged_kv_last_page_len <= page_size
>>> paged_kv_last_page_len = torch.tensor(
...     [1, 7, 14, 4, 3, 1, 16], dtype=torch.int32, device="cuda:0"
... )
>>> q_at_layer = torch.randn(num_layers, nnz_qo, num_qo_heads, head_dim).half().to("cuda:0")
>>> kv_cache_at_layer = torch.randn(
...     num_layers, max_num_pages, 2, page_size, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
... )
>>> # create auxiliary data structures for batch prefill attention
>>> prefill_wrapper.plan(
...     qo_indptr,
...     paged_kv_indptr,
...     paged_kv_indices,
...     paged_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     causal=True,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     kv_cache = kv_cache_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o = prefill_wrapper.run(q, kv_cache)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([100, 64, 128])
>>>
>>> # below is another example of creating custom mask for batch prefill attention
>>> mask_arr = []
>>> qo_len = (qo_indptr[1:] - qo_indptr[:-1]).cpu().tolist()
>>> kv_len = (page_size * (paged_kv_indptr[1:] - paged_kv_indptr[:-1] - 1) + paged_kv_last_page_len).cpu().tolist()
>>> for i in range(batch_size):
...     mask_i = torch.tril(
...         torch.full((qo_len[i], kv_len[i]), True, device="cuda:0"),
...         diagonal=(kv_len[i] - qo_len[i]),
...     )
...     mask_arr.append(mask_i.flatten())
...
>>> mask = torch.cat(mask_arr, dim=0)
>>> prefill_wrapper.plan(
...     qo_indptr,
...     paged_kv_indptr,
...     paged_kv_indices,
...     paged_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     custom_mask=mask,
... )
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     kv_cache = kv_cache_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o_custom = prefill_wrapper.run(q, kv_cache)
...     assert torch.allclose(o_custom, outputs[i], rtol=1e-3, atol=1e-3)
...

Note

To accelerate computation, FlashInfer’s batch prefill/append attention operators create some auxiliary data structures, these data structures can be reused across multiple prefill/append attention calls (e.g. different Transformer layers). This wrapper class manages the lifecycle of these data structures.

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None, variant_owns_mask: bool = False) → None

Constructor of BatchPrefillWithPagedKVCacheWrapper.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The user reserved workspace buffer used to store intermediate attention results in split-k algorithm. The recommended size is 128MB, the device of the workspace buffer should be the same as the device of the input tensors. The buffer must be 16-byte aligned; tensors created by torch.empty satisfy this on supported devices.

  • kv_layout (str) – The layout of the input k/v tensors, could be either NHD or HND.

  • use_cuda_graph (bool) – Whether to enable CUDA graph capture for the prefill kernels, if enabled, the auxiliary data structures will be stored in provided buffers. The batch_size cannot change during the lifecycle of this wrapper when CUDAGraph is enabled.

  • qo_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the qo_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • paged_kv_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_indptr array, the size of this buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • paged_kv_indices_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_indices array, should be large enough to store the maximum possible size of the paged_kv_indices array during the lifetime of the wrapper. This argument is only effective when use_cuda_graph is True.

  • paged_kv_last_page_len_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_last_page_len array, the size of the buffer should be [batch_size]. This argument is only effective when use_cuda_graph is True.

  • custom_mask_buf (Optional[torch.Tensor]) – The user reserved buffer to store the custom mask tensor, should be large enough to store the maximum possible size of the packed custom mask tensor during the lifetime of the wrapper. This argument is only effective when use_cuda_graph is set to True and the custom mask will be used in attention computation.

  • mask_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the mask_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True and the custom mask will be used in attention computation.

  • backend (str) – The implementation backend, could be auto/fa2/fa3/cudnn/trtllm-gen cute-dsl/cute-dsl-prims. Defaults to auto. If set to auto, the wrapper will automatically choose the backend based on the device architecture and kernel availability. The cute-dsl backend uses the CuTe DSL attention kernel for Blackwell (SM100+). cute-dsl-prims backend is a SM120-only FP8 attention backend implemented using CUTLASS primitives.

  • jit_args (Optional[List[Any]]) – If provided, the wrapper will use the provided arguments to create the JIT module, otherwise, the wrapper will use default attention implementation.

  • jit_kwargs (Optional[Dict[str, Any]]) – The keyword arguments to create the JIT module, defaults to None.

  • variant_owns_mask (bool) – If True, the attention variant supplied through jit_args computes the complete attention mask in its LogitsMask hook, and MaskMode.CUSTOM is selected without a mask tensor: the kernel evaluates LogitsMask on every KV tile instead of only on the causal/window boundary tiles, and no custom_mask/packed_custom_mask needs to be passed to plan(). Requires jit_args, because the default attention variant dereferences the custom mask buffer under MaskMode.CUSTOM. Only supported with backend="fa2" (the SM90 batch prefill kernels reject MaskMode.CUSTOM), and incompatible with prefix_len_ptr (multi-item scoring), which selects a different mask mode. Defaults to False.

plan(qo_indptr: Tensor, paged_kv_indptr: Tensor, paged_kv_indices: Tensor, paged_kv_last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, page_size: int, head_dim_vo: int | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, sm_scale: float | None = None, window_left: int = -1, logits_soft_cap: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, non_blocking: bool = True, prefix_len_ptr: Tensor | None = None, token_pos_in_items_ptr: Tensor | None = None, token_pos_in_items_len: int = 0, max_item_len_ptr: Tensor | None = None, seq_lens: Tensor | None = None, seq_lens_q: Tensor | None = None, block_tables: Tensor | None = None, max_token_per_sequence: int | None = None, max_sequence_kv: int | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False) → None

Plan batch prefill/append attention on Paged KV-Cache for given problem specification.

Parameters:
  • qo_indptr (torch.Tensor) – The indptr of the query/output tensor, shape: [batch_size + 1]. For the cudnn backend this is interpreted in element units (cumsum(seq_lens_q) * num_qo_heads * head_dim_qk), not token units.

  • paged_kv_indptr (torch.Tensor) – The indptr of the paged kv-cache, shape: [batch_size + 1].

  • paged_kv_indices (torch.Tensor) – The page indices of the paged kv-cache, shape: [paged_kv_indptr[-1]].

  • paged_kv_last_page_len (torch.Tensor) – The number of entries in the last page of each request in the paged kv-cache, shape: [batch_size].

  • num_qo_heads (int) – The number of query/output heads.

  • num_kv_heads (int) – The number of key/value heads.

  • head_dim_qk (int) – The dimension of the query/key heads.

  • page_size (int) – The size of each page in the paged kv-cache.

  • head_dim_vo (Optional[int]) – The dimension of the value/output heads, if not provided, will be set to head_dim_qk.

  • custom_mask (Optional[torch.Tensor]) –

    The flattened boolean mask tensor, shape: (sum(q_len[i] * k_len[i] for i in range(batch_size)). The elements in the mask tensor should be either True or False, where False means the corresponding element in the attention matrix will be masked out.

    Please refer to the mask layout for more details about flattened layout of mask tensor.

    When custom_mask is provided, and packed_custom_mask is not, the function will pack the custom mask tensor into a 1D packed mask tensor, which introduces additional overhead.

  • packed_custom_mask (Optional[torch.Tensor]) – The 1D packed uint8 mask tensor, if provided, the custom_mask will be ignored. The packed mask tensor is generated by flashinfer.quantization.packbits().

  • causal (bool) – Whether to apply causal mask to the attention matrix. This is only effective when custom_mask is not provided in plan().

  • pos_encoding_mode (str) – The position encoding applied inside attention kernels, could be NONE/ROPE_LLAMA (LLAMA style rotary embedding) /ALIBI. Default is NONE.

  • use_fp16_qk_reduction (bool) – Whether to use f16 for qk reduction (faster at the cost of slight precision loss).

  • window_left (int) – The left (inclusive) window size for the attention window, when set to -1, the window size will be set to the full length of the sequence. Defaults to -1.

  • logits_soft_cap (Optional[float]) – The attention logits soft capping value (used in Gemini, Grok and Gemma-2, etc.), if not provided, will be set to 0. If greater than 0, the logits will be capped according to formula: \(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\), where \(x\) is the input logits.

  • sm_scale (Optional[float]) – The scale used in softmax, if not provided, will be set to 1.0 / sqrt(head_dim).

  • rope_scale (Optional[float]) – The scale used in RoPE interpolation, if not provided, will be set to 1.0.

  • rope_theta (Optional[float]) – The theta used in RoPE, if not provided, will be set to 1e4.

  • q_data_type (Union[str, torch.dtype]) – The data type of the query tensor, defaults torch.float16.

  • kv_data_type (Optional[Union[str, torch.dtype]]) – The data type of the key/value tensor. If None, will be set to q_data_type.

  • o_data_type (Optional[Union[str, torch.dtype]]) – The data type of the output tensor. If None, will be set to q_data_type. For FP8 inputs, this should typically be set to torch.float16 or torch.bfloat16.

  • non_blocking (bool) – Whether to copy the input tensors to the device asynchronously, defaults to True.

  • prefix_len_ptr (Optional[torch.Tensor]) – prefix length. A uint32 1D tensor indicating the prefix length of each prompt. The tensor size is equal to the batch size.

  • token_pos_in_items_ptr (Optional[torch.Tensor]) – A uint16 1D tensor (it will be converted to uint16 in flashinfer) indicating the token position of each item and started from 0 (delimiter) for each item. E.g., if we have 3 items of length 3, 2, 4 respectively for this member. This vector will be looking like [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0] with 4 delimiters indexed as 0. For batch size > 1, we will concat them as 1D with zero paddings to make sure each has the same length, the padding length is defined by token_pos_in_items_len - length of the raw token_pos_in_items_ptr for each prompt.

  • token_pos_in_items_len (int) – zero padding length for token_pos_in_items_ptr to better handle the bsz > 1 case. Still using the above 3,2,4 example. If we set token_pos_in_items_len to be 20, it will be [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0] with 7 padded zeros. (note there’re 8 zeros in the end where the first one is the delimiter token 0 in the end of the prompt)

  • max_item_len_ptr (Optional[torch.Tensor]) – a uint16 vector contains the max token length of all items for each prompt

  • seq_lens (Optional[torch.Tensor]) – A uint32 1D tensor indicating the kv sequence length of each prompt. shape: [batch_size].

  • seq_lens_q (Optional[torch.Tensor]) – A uint32 1D tensor indicating the q sequence length of each prompt. shape: [batch_size]. If not provided, will be set to the same value as seq_lens.

  • block_tables (Optional[torch.Tensor]) – A uint32 2D tensor indicating the block table of each prompt. shape: [batch_size, max_num_blocks_per_seq].

  • max_token_per_sequence (Optional[int],) – Required for cudnn backend. This is the scalar max token length of each sequence.

  • max_sequence_kv (Optional[int],) – Required for cudnn backend. This is the scalar max sequence length of each sequence in kv cache.

  • fixed_split_size (Optional[int],) – The fixed split size for FA2 split-kv prefill/decode in pages. Recommend setting to the average sequence length of your workload. When enabled, will lead to deterministic softmax score reduction in the merge_states kernel, and therefore batch-size invariant outputs. See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ Note that compatibility with CUDA graph is NOT guaranteed, as even when bs is fixed, kv seq len can change and lead to a varied number of launched CTAs.

  • disable_split_kv (bool,) – Whether to disable the split-kv for determinism in CUDA Graph, defaults to False.

Note

The plan() method should be called before any run() or run_return_lse() calls, auxiliary data structures will be created during this call and cached for multiple kernel runs.

The num_qo_heads must be a multiple of num_kv_heads. If num_qo_heads is not equal to num_kv_heads, the function will use grouped query attention.

The plan() method cannot be used in Cuda Graph or in torch.compile.

prewarm_paged_kv_stride_variant(variant: str = 'independent') → None

Load a lazy paged-KV-stride variant after plan().

Call this method before CUDA graph capture when a planned standard FA2 wrapper will receive K and V tensors with different data strides.

Parameters:

variant (str) – The paged-KV-stride variant to prewarm. The only supported value is "independent" (the default), for K and V with different data strides.

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) → None

Reset the workspace buffer.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The new float workspace buffer, the device of the new float workspace buffer should be the same as the device of the input tensors.

  • int_workspace_buffer (torch.Tensor) – The new int workspace buffer, the device of the new int workspace buffer should be the same as the device of the input tensors.

run(q: Tensor, paged_kv_cache: torch.Tensor | Tuple[torch.Tensor, torch.Tensor], *args, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None, window_left: int | None = None, sinks: Tensor | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None, skip_softmax_threshold_scale_factor: float | None = None, use_fp16_softmax: bool | None = None, uses_spcompress: bool | None = None) → Tensor
run(q: Tensor, paged_kv_cache: torch.Tensor | Tuple[torch.Tensor, torch.Tensor], *args, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None, window_left: int | None = None, sinks: Tensor | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None, skip_softmax_threshold_scale_factor: float | None = None, use_fp16_softmax: bool | None = None, uses_spcompress: bool | None = None) → Tuple[Tensor, Tensor]

Compute batch prefill/append attention between query and paged kv-cache.

Parameters:
  • q (torch.Tensor) – The query tensor, shape: [qo_indptr[-1], num_qo_heads, head_dim]

  • paged_kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) –

    The paged KV-Cache stored as a tuple of tensors or a single tensor:

    • a tuple (k_cache, v_cache) of 4-D tensors, each with shape: [max_num_pages, page_size, num_kv_heads, head_dim] if kv_layout is NHD, and [max_num_pages, num_kv_heads, page_size, head_dim] if kv_layout is HND.

    • a single 5-D tensor with shape: [max_num_pages, 2, page_size, num_kv_heads, head_dim] if kv_layout is NHD, and [max_num_pages, 2, num_kv_heads, page_size, head_dim] if kv_layout is HND. Where paged_kv_cache[:, 0] is the key-cache and paged_kv_cache[:, 1] is the value-cache.

  • *args – Additional arguments for custom kernels.

  • q_scale (Optional[Union[float, torch.Tensor]]) – The calibration scale of query for fp8 input, if not provided, will be set to 1.0.

  • k_scale (Optional[Union[float, torch.Tensor]]) – The calibration scale of key for fp8 or nvfp4 input, if not provided, will be set to 1.0.

  • v_scale (Optional[Union[float, torch.Tensor]]) – The calibration scale of value for fp8 or nvfp4 input, if not provided, will be set to 1.0.

  • out (Optional[torch.Tensor]) – The output tensor, if not provided, will be allocated internally.

  • lse (Optional[torch.Tensor]) – The log-sum-exp of attention logits, if not provided, will be allocated internally.

  • return_lse (bool) – Whether to return the logsumexp of attention output

  • enable_pdl (bool) – Whether to enable Programmatic Dependent Launch (PDL). See https://docs.nvidia.com/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization Only effective on backends and devices that support PDL.

  • window_left (Optional[int]) – Per-call override for the left (inclusive) sliding-window size. When None, the value supplied to plan() is used. Pass -1 to disable the sliding window for this call.

  • sinks (Optional[torch.Tensor]) – Per-head attention-sink logits. When provided, the kernel applies the attention-with-sink variant: an additional virtual token whose logit is sinks[head_idx] is appended to each row of the softmax denominator. Shape: [num_qo_heads], dtype float32.

  • kv_cache_sf (Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]]) –

    Per-block scale factors for NVFP4 KV cache. Accepts the same formats as paged_kv_cache:

    • a tuple (k_scales, v_scales) of 4-D tensors, each with shape: [num_pages, page_size, num_kv_heads, head_dim // 16] if kv_layout is NHD, and [num_pages, num_kv_heads, page_size, head_dim // 16] if kv_layout is HND.

    • a single 5-D tensor with shape: [num_pages, 2, page_size, num_kv_heads, head_dim // 16] if kv_layout is NHD, and [num_pages, 2, num_kv_heads, page_size, head_dim // 16] if kv_layout is HND, where dim 1 holds k (index 0) and v (index 1) scales.

    Both tensors have dtype torch.float8_e4m3fn. k_scales uses a linear (row-major) layout, while v_scales must use TRT-LLM’s 4-token interleaved layout within each [page_size, head_dim // 16] tile if backend is trtllm-gen. Use flashinfer.fp4_quantization.nvfp4_quantize_paged_kv_cache() to produce correctly formatted scale factors.

    For the trtllm-gen backend with NHD layout, scale tensors are transposed to HND internally (incurring a copy). Use HND for better performance. Currently, NVFP4 KV supports fa2 and trtllm-gen backend.

  • use_fp16_softmax (Optional[bool]) – trtllm-gen backend only. Select the …Fp16Softmax… cubin variant (FP16 softmax accumulator). Currently only shipped for BF16 Q/KV/O context kernels. Ignored by other backends.

  • uses_spcompress (Optional[bool]) – trtllm-gen backend only. Select the …Spcomp… cubin variant (sparse compression). Currently only shipped for FP8 Q context kernels. Ignored by other backends.

  • skip_softmax_threshold_scale_factor (Optional[float]) – Threshold scale factor for skipping softmax operations. Providing a value enables skip-softmax sparsity as described in https://arxiv.org/abs/2512.12087. Defaults to None (standard attention). Higher values yield faster kernels at the cost of accuracy; the effective threshold equals the supplied factor divided by the context length.

Returns:

If return_lse is False, the attention output, shape: [qo_indptr[-1], num_qo_heads, head_dim]. If return_lse is True, a tuple of two tensors:

  • The attention output, shape: [qo_indptr[-1], num_qo_heads, head_dim].

  • The logsumexp of attention output, shape: [qo_indptr[-1], num_qo_heads].

Return type:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

workspace_size(qo_indptr: Tensor, paged_kv_indptr: Tensor, paged_kv_indices: Tensor, paged_kv_last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, page_size: int, head_dim_vo: int | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, sm_scale: float | None = None, window_left: int = -1, logits_soft_cap: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, prefix_len_ptr: Tensor | None = None, token_pos_in_items_ptr: Tensor | None = None, token_pos_in_items_len: int = 0, max_item_len_ptr: Tensor | None = None, seq_lens: Tensor | None = None, seq_lens_q: Tensor | None = None, block_tables: Tensor | None = None, max_token_per_sequence: int | None = None, max_sequence_kv: int | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False) → Tuple[int, int]

Return the caller-owned workspace size required by plan().

The returned tuple is (float_workspace_size, int_workspace_size) in bytes. The inputs follow plan(); host-side planning tensors such as qo_indptr and paged_kv_indptr are copied to CPU in the same way as plan(). The wrapper’s float workspace buffer is only used as the device/stream selector for the underlying module query. This method does not allocate buffers and does not mutate cached plan state.

Parameters:
  • qo_indptr (torch.Tensor) – The indptr of the query/output tensor, shape: [batch_size + 1].

  • paged_kv_indptr (torch.Tensor) – The indptr of the paged kv-cache, shape: [batch_size + 1].

  • paged_kv_indices (torch.Tensor) – The page indices of the paged kv-cache, shape: [paged_kv_indptr[-1]].

  • paged_kv_last_page_len (torch.Tensor) – The number of entries in the last page of each request in the paged kv-cache, shape: [batch_size].

  • num_qo_heads (int) – The number of query/output heads.

  • num_kv_heads (int) – The number of key/value heads.

  • head_dim_qk (int) – The dimension of the query/key heads.

  • page_size (int) – The size of each page in the paged kv-cache.

  • head_dim_vo (Optional[int]) – The dimension of the value/output heads. If not provided, defaults to head_dim_qk.

  • custom_mask (Optional[torch.Tensor]) – The flattened boolean mask tensor; when provided the packed mask is generated internally. See plan() for the full description.

  • packed_custom_mask (Optional[torch.Tensor]) – The 1D packed uint8 mask tensor. If provided, custom_mask is ignored.

  • causal (bool) – Whether to apply a causal mask. Defaults to False.

  • pos_encoding_mode (str) – The position encoding applied inside attention kernels, could be NONE/ROPE_LLAMA (LLAMA style rotary embedding) /ALIBI. Defaults to NONE.

  • use_fp16_qk_reduction (bool) – Whether to use fp16 for qk reduction. Defaults to False.

  • sm_scale (Optional[float]) – Softmax scale. If None, defaults to 1.0 / sqrt(head_dim_qk).

  • window_left (int) – The left (inclusive) window size for the attention window, when set to -1, the window size will be set to the full length of the sequence. Defaults to -1.

  • logits_soft_cap (Optional[float]) – The attention logits soft capping value (used in Gemini, Grok and Gemma-2, etc.), if not provided, will be set to 0.

  • rope_scale (Optional[float]) – Scale factor applied during RoPE interpolation. Defaults to 1.0 when None.

  • rope_theta (Optional[float]) – Base value for the RoPE frequencies. Defaults to 1e4 when None.

  • q_data_type (Union[str, torch.dtype]) – The data type of the query tensor. Defaults to torch.float16.

  • kv_data_type (Optional[Union[str, torch.dtype]]) – The data type of the key/value tensor. If None, will be set to q_data_type.

  • o_data_type (Optional[Union[str, torch.dtype]]) – The data type of the output tensor. If None, will be set to q_data_type.

  • prefix_len_ptr (Optional[torch.Tensor]) – A uint32 1D tensor indicating the prefix length of each prompt, shape: [batch_size].

  • token_pos_in_items_ptr (Optional[torch.Tensor]) – A uint16 1D tensor indicating the token position of each item within its item.

  • token_pos_in_items_len (int) – Zero-padding length for token_pos_in_items_ptr. Defaults to 0.

  • max_item_len_ptr (Optional[torch.Tensor]) – A uint16 vector containing the max token length of all items for each prompt.

  • seq_lens (Optional[torch.Tensor]) – A uint32 1D tensor indicating the kv sequence length of each prompt, shape: [batch_size].

  • seq_lens_q (Optional[torch.Tensor]) – A uint32 1D tensor indicating the q sequence length of each prompt, shape: [batch_size].

  • block_tables (Optional[torch.Tensor]) – A uint32 2D tensor indicating the block table of each prompt, shape: [batch_size, max_num_blocks_per_seq].

  • max_token_per_sequence (Optional[int]) – Required for cuDNN backend. The scalar max token length of each sequence.

  • max_sequence_kv (Optional[int]) – Maximum number of KV tokens per sequence for cuDNN backend.

  • fixed_split_size (Optional[int]) – The fixed split size for split-kv prefill, in pages.

  • disable_split_kv (bool) – Whether to disable the split-kv. Defaults to False.

Returns:

(float_workspace_size, int_workspace_size) in bytes.

Return type:

Tuple[int, int]

Example

>>> float_bytes, int_bytes = wrapper.workspace_size(...)
>>> wrapper.reset_workspace_buffer(
...     torch.empty(float_bytes, dtype=torch.uint8, device="cuda"),
...     torch.empty(int_bytes, dtype=torch.uint8, device="cuda"),
... )
>>> wrapper.plan(...)
class flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, kv_indptr_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None, variant_owns_mask: bool = False)

Wrapper class for prefill/append attention with ragged (tensor) kv-cache for batch of requests.

Check our tutorial for ragged kv-cache layout.

Example

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 16
>>> head_dim = 128
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> prefill_wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> nnz_kv = 100
>>> nnz_qo = 100
>>> qo_indptr = torch.tensor(
...     [0, 33, 44, 55, 66, 77, 88, nnz_qo], dtype=torch.int32, device="cuda:0"
... )
>>> kv_indptr = qo_indptr.clone()
>>> q_at_layer = torch.randn(num_layers, nnz_qo, num_qo_heads, head_dim).half().to("cuda:0")
>>> k_at_layer = torch.randn(num_layers, nnz_kv, num_kv_heads, head_dim).half().to("cuda:0")
>>> v_at_layer = torch.randn(num_layers, nnz_kv, num_kv_heads, head_dim).half().to("cuda:0")
>>> # create auxiliary data structures for batch prefill attention
>>> prefill_wrapper.plan(
...     qo_indptr,
...     kv_indptr,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     causal=True,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     k = k_at_layer[i]
...     v = v_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o = prefill_wrapper.run(q, k, v)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([100, 64, 128])
>>>
>>> # below is another example of creating custom mask for batch prefill attention
>>> mask_arr = []
>>> qo_len = (qo_indptr[1:] - qo_indptr[:-1]).cpu().tolist()
>>> kv_len = (kv_indptr[1:] - kv_indptr[:-1]).cpu().tolist()
>>> for i in range(batch_size):
...     mask_i = torch.tril(
...         torch.full((qo_len[i], kv_len[i]), True, device="cuda:0"),
...         diagonal=(kv_len[i] - qo_len[i]),
...     )
...     mask_arr.append(mask_i.flatten())
...
>>> mask = torch.cat(mask_arr, dim=0)
>>> prefill_wrapper.plan(
...     qo_indptr,
...     kv_indptr,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     custom_mask=mask
... )
>>> outputs_custom_mask = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     k = k_at_layer[i]
...     v = v_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o_custom = prefill_wrapper.run(q, k, v)
...     assert torch.allclose(o_custom, outputs[i], rtol=1e-3, atol=1e-3)
...
>>> outputs_custom_mask[0].shape
torch.Size([100, 64, 128])

Note

To accelerate computation, FlashInfer’s batch prefill/append attention operators create some auxiliary data structures, these data structures can be reused across multiple prefill/append attention calls (e.g. different Transformer layers). This wrapper class manages the lifecycle of these data structures.

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, kv_indptr_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None, variant_owns_mask: bool = False) → None

Constructor of BatchPrefillWithRaggedKVCacheWrapper.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The user reserved float workspace buffer used to store intermediate attention results in the split-k algorithm. The recommended size is 128MB, the device of the workspace buffer should be the same as the device of the input tensors. The buffer must be 16-byte aligned; tensors created by torch.empty satisfy this on supported devices.

  • kv_layout (str) – The layout of the input k/v tensors, could be either NHD or HND.

  • use_cuda_graph (bool) – Whether to enable CUDA graph capture for the prefill kernels, if enabled, the auxiliary data structures will be stored as the provided buffers.

  • qo_indptr_buf (Optional[torch.Tensor]) – The user reserved GPU buffer to store the qo_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • kv_indptr_buf (Optional[torch.Tensor]) – The user reserved GPU buffer to store the kv_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • custom_mask_buf (Optional[torch.Tensor]) – The user reserved GPU buffer to store the custom mask tensor, should be large enough to store the maximum possible size of the packed custom mask tensor during the lifetime of the wrapper. This argument is only effective when use_cuda_graph is True and custom mask will be used in attention computation.

  • mask_indptr_buf (Optional[torch.Tensor]) – The user reserved GPU buffer to store the mask_indptr array, the size of the buffer should be [batch_size]. This argument is only effective when use_cuda_graph is True and custom mask will be used in attention computation.

  • backend (str) – The implementation backend, could be auto/fa2/fa3/cudnn/cutlass or cute-dsl/cute-dsl-prims/cutile. Defaults to auto. If set to auto, the wrapper will automatically choose the backend based on the device architecture and kernel availability. The cute-dsl backend uses the CuTe DSL attention kernel for Blackwell (SM100+). cute-dsl-prims is an explicit SM120-only packed FP8 prefill backend. The cutile backend uses the pure cuda.tile Python prefill kernel (Blackwell, opt-in); it requires qo_indptr == kv_indptr (equal-length prefill).

  • jit_args (Optional[List[Any]]) – If provided, the wrapper will use the provided arguments to create the JIT module, otherwise, the wrapper will use default attention implementation.

  • jit_kwargs (Optional[Dict[str, Any]]) – The keyword arguments to create the JIT module, defaults to None.

  • variant_owns_mask (bool) – If True, the attention variant supplied through jit_args computes the complete attention mask in its LogitsMask hook, and MaskMode.CUSTOM is selected without a mask tensor: the kernel evaluates LogitsMask on every KV tile instead of only on the causal/window boundary tiles, and no custom_mask/packed_custom_mask needs to be passed to plan(). Requires jit_args, because the default attention variant dereferences the custom mask buffer under MaskMode.CUSTOM. Only supported with backend="fa2" (the SM90 batch prefill kernels reject MaskMode.CUSTOM), and incompatible with prefix_len_ptr (multi-item scoring), which selects a different mask mode. Defaults to False.

plan(qo_indptr: Tensor, kv_indptr: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, head_dim_vo: int | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, window_left: int = -1, logits_soft_cap: float | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, non_blocking: bool = True, prefix_len_ptr: Tensor | None = None, token_pos_in_items_ptr: Tensor | None = None, token_pos_in_items_len: int = 0, max_item_len_ptr: Tensor | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False, seq_lens: Tensor | None = None, seq_lens_q: Tensor | None = None, max_token_per_sequence: int | None = None, max_sequence_kv: int | None = None, v_indptr: Tensor | None = None, o_indptr: Tensor | None = None) → None

Plan batch prefill/append attention on Ragged KV-Cache for given problem specification.

Parameters:
  • qo_indptr (torch.Tensor) – The indptr of the query/output tensor, shape: [batch_size + 1]. For the cudnn backend the qo_indptr and kv_indptr are interpreted in element units (cumsum(seq_lens) * num_heads * head_dim_qk), not token units. The cudnn backend also requires kv_layout="NHD".

  • kv_indptr (torch.Tensor) – The indptr of the key/value tensor, shape: [batch_size + 1].

  • num_qo_heads (int) – The number of query/output heads.

  • num_kv_heads (int) – The number of key/value heads.

  • head_dim_qk (int) – The dimension of the heads on query/key tensor.

  • head_dim_vo (Optional[int]) – The dimension of the heads on value/output tensor. If not provided, will be set to head_dim_qk.

  • custom_mask (Optional[torch.Tensor]) –

    The flattened boolean mask tensor, shape: (sum(q_len[i] * k_len[i] for i in range(batch_size)). The elements in the mask tensor should be either True or False, where False means the corresponding element in the attention matrix will be masked out.

    Please refer to the mask layout for more details about flattened layout of mask tensor.

    When custom_mask is provided, and packed_custom_mask is not, the function will pack the custom mask tensor into a 1D packed mask tensor, which introduces additional overhead.

  • packed_custom_mask (Optional[torch.Tensor]) –

    The 1D packed uint8 mask tensor, if provided, the custom_mask will be ignored. The packed mask tensor is generated by flashinfer.quantization.packbits().

    If provided, the custom mask will be added to the attention matrix before softmax and after scaling. The mask tensor should be in the same device as the input tensors.

  • causal (bool) – Whether to apply causal mask to the attention matrix. This argument is ignored if mask is provided in plan().

  • pos_encoding_mode (str) – The position encoding applied inside attention kernels, could be NONE/ROPE_LLAMA (LLAMA style rotary embedding) /ALIBI. Default is NONE.

  • use_fp16_qk_reduction (bool) – Whether to use f16 for qk reduction (faster at the cost of slight precision loss).

  • window_left (int) – The left (inclusive) window size for the attention window, when set to -1, the window size will be set to the full length of the sequence. Defaults to -1.

  • logits_soft_cap (Optional[float]) – The attention logits soft capping value (used in Gemini, Grok and Gemma-2, etc.), if not provided, will be set to 0. If greater than 0, the logits will be capped according to formula: \(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\), where \(x\) is the input logits.

  • sm_scale (Optional[float]) – The scale used in softmax, if not provided, will be set to 1.0 / sqrt(head_dim_qk).

  • rope_scale (Optional[float]) – The scale used in RoPE interpolation, if not provided, will be set to 1.0.

  • rope_theta (Optional[float]) – The theta used in RoPE, if not provided, will be set to 1e4.

  • q_data_type (Union[str, torch.dtype]) – The data type of the query tensor, defaults to torch.float16.

  • kv_data_type (Optional[Union[str, torch.dtype]]) – The data type of the key/value tensor. If None, will be set to q_data_type.

  • o_data_type (Optional[Union[str, torch.dtype]]) – The data type of the output tensor. If None, will be set to q_data_type. For FP8 inputs, this should typically be set to torch.float16 or torch.bfloat16.

  • non_blocking (bool) – Whether to copy the input tensors to the device asynchronously, defaults to True.

  • prefix_len_ptr (Optional[torch.Tensor]) – prefix length. A uint32 1D tensor indicating the prefix length of each prompt. The tensor size is equal to the batch size.

  • token_pos_in_items_ptr (Optional[torch.Tensor]) – A uint16 1D tensor (it will be converted to uint16 in flashinfer) indicating the token position of each item and started from 0 (delimiter) for each item. E.g., if we have 3 items of length 3, 2, 4 respectively for this member. This vector will be looking like [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0] with 4 delimiters indexed as 0. For batch size > 1, we will concat them as 1D with zero paddings to make sure each has the same length, the padding length is defined by token_pos_in_items_len - length of the raw token_pos_in_items_ptr for each prompt.

  • token_pos_in_items_len (int) – zero padding length for token_pos_in_items_ptr to better handle the bsz > 1 case. Still using the above 3,2,4 example. If we set token_pos_in_items_len to be 20, it will be [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0] with 7 padded zeros. (note there’re 8 zeros in the end where the first one is the delimiter token 0 in the end of the prompt)

  • max_item_len_ptr (Optional[torch.Tensor]) – a uint16 vector contains the max token length of all items for each prompt

  • fixed_split_size (Optional[int],) – The fixed split size for split-kv FA2 prefill/decode, in pages. Recommend setting to the average sequence length of your workload. When enabled, will lead to deterministic softmax score reduction in the merge_states kernel, and therefore batch-size invariant outputs. See https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ Note that compatibility with CUDA graph is NOT guaranteed, as even when bs is fixed, kv seq len can change and lead to a varied number of launched CTAs.

  • disable_split_kv (bool,) – Whether to disable the split-kv for determinism in CUDA Graph, defaults to False.

  • seq_lens (Optional[torch.Tensor]) – A uint32 1D tensor indicating the kv sequence length of each prompt. shape: [batch_size].

  • seq_lens_q (Optional[torch.Tensor]) – A uint32 1D tensor indicating the q sequence length of each prompt. shape: [batch_size]. If not provided, will be set to the same value as seq_lens.

  • max_token_per_sequence (Optional[int],) – Required for cudnn backend. This is the scalar max token length of each sequence.

  • max_sequence_kv (Optional[int],) – Required for cudnn backend. This is the scalar max sequence length of each sequence in kv cache.

  • v_indptr (Optional[torch.Tensor]) – Required for cudnn backend. This is the indptr of the value tensor.

  • o_indptr (Optional[torch.Tensor]) – Required for cudnn backend. This is the indptr of the output tensor.

Note

The plan() method should be called before any run() or run_return_lse() calls, auxiliary data structures will be created during this plan call and cached for multiple kernel runs.

The num_qo_heads must be a multiple of num_kv_heads. If num_qo_heads is not equal to num_kv_heads, the function will use grouped query attention.

The plan() method cannot be used in Cuda Graph or in torch.compile.

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer) → None

Reset the workspace buffer.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The new float workspace buffer, the device of the new float workspace buffer should be the same as the device of the input tensors.

  • int_workspace_buffer (torch.Tensor) – The new int workspace buffer, the device of the new int workspace buffer should be the same as the device of the input tensors.

run(q: Tensor, k: Tensor, v: Tensor, *args, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None) → Tensor
run(q: Tensor, k: Tensor, v: Tensor, *args, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None, kv_cache_sf: torch.Tensor | Tuple[torch.Tensor, torch.Tensor] | None = None) → Tuple[Tensor, Tensor]

Compute batch prefill/append attention between query and kv-cache stored as ragged tensor.

Parameters:
  • q (torch.Tensor) – The query tensor, shape: [qo_indptr[-1], num_qo_heads, head_dim_qk]

  • k (torch.Tensor) – The key tensor, shape: [kv_indptr[-1], num_kv_heads, head_dim_qk]

  • v (torch.Tensor) – The value tensor, shape: [kv_indptr[-1], num_kv_heads, head_dim_vo]

  • *args – Additional arguments for the custom kernel.

  • q_scale (Optional[float]) – The calibration scale of fp8 query, if not provided, will be set to 1.0.

  • k_scale (Optional[float]) – The calibration scale of fp8 or nvfp4 key, if not provided, will be set to 1.0.

  • v_scale (Optional[float]) – The calibration scale of fp8 or nvfp4 value, if not provided, will be set to 1.0.

  • o_scale (Optional[float]) – The calibration scale of output, if not provided, will be set to 1.0.

  • out (Optional[torch.Tensor]) – The output tensor, if not provided, will be allocated internally.

  • lse (Optional[torch.Tensor]) – The log-sum-exp of attention logits, if not provided, will be allocated internally.

  • return_lse (bool) – Whether to return the logsumexp of attention output

  • enable_pdl (bool) – Whether to enable Programmatic Dependent Launch (PDL). See https://docs.nvidia.com/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization Only effective on backends and devices that support PDL.

  • kv_cache_sf (Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]]) – Per-block scale factors for NVFP4 KV input. Accepts either a single packed scale tensor or a (k_scales, v_scales) tuple matching the structure expected by the chosen backend. When None (default), the kernel runs without NVFP4 KV scaling. See flashinfer.fp4_quantization.nvfp4_quantize() for layout details.

Returns:

If return_lse is False, the attention output, shape: [qo_indptr[-1], num_qo_heads, head_dim_vo]. If return_lse is True, a tuple of two tensors:

  • The attention output, shape: [qo_indptr[-1], num_qo_heads, head_dim_vo].

  • The logsumexp of attention output, shape: [qo_indptr[-1], num_qo_heads].

Return type:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

Unified BatchAttention

The BatchAttention class provides a holistic attention wrapper that automatically dispatches between paged-prefill and paged-decode based on per-request sequence lengths. It is the recommended entry point for serving stacks that batch mixed prefill/decode requests in a single kernel launch.

class flashinfer.attention.BatchAttention(kv_layout: str = 'NHD', device: str = 'cuda')

Holistic batched attention wrapper that fuses paged-prefill and paged-decode requests into a single kernel launch.

BatchAttention dispatches between prefill-style and decode-style execution per request based on the qo_indptr / kv_indptr ranges supplied to plan(), so a serving stack can submit a mixed batch (e.g. some prompts in prefill, others in decode) without splitting it into two separate wrappers. Workspace buffers are owned by the instance and reused across plan() / run() calls.

Parameters:
  • kv_layout (str) – Layout of the paged KV-cache tensors, either "NHD" (token-major) or "HND" (head-major). Defaults to "NHD".

  • device (str) – CUDA device that owns the internal workspace buffers, e.g. "cuda" or "cuda:0". Defaults to "cuda".

__init__(kv_layout: str = 'NHD', device: str = 'cuda')

Allocate workspace buffers and bind the wrapper to a CUDA device.

See BatchAttention for the meaning of each parameter.

plan(qo_indptr: Tensor, kv_indptr: Tensor, kv_indices: Tensor, kv_len_arr: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, head_dim_vo: int, page_size: int, causal: bool = False, sm_scale: float | None = None, logits_soft_cap: float | None = None, q_data_type: dtype = torch.bfloat16, kv_data_type: dtype = torch.bfloat16, use_profiler: bool = False) → None

Plan the holistic attention kernel for a specific batch shape.

Should be called before any run() call. The plan is cached on the instance and reused across subsequent run() invocations with the same layout.

Parameters:
  • qo_indptr (torch.Tensor) – CSR-style query offsets, shape [batch_size + 1], dtype int32.

  • kv_indptr (torch.Tensor) – CSR-style page offsets into kv_indices, shape [batch_size + 1], dtype int32.

  • kv_indices (torch.Tensor) – Page indices into the paged KV-cache, shape [kv_indptr[-1]], dtype int32.

  • kv_len_arr (torch.Tensor) – Per-request KV-cache lengths in tokens, shape [batch_size], dtype int32.

  • num_qo_heads (int) – Number of query / output heads.

  • num_kv_heads (int) – Number of key / value heads. Must divide num_qo_heads.

  • head_dim_qk (int) – Per-head dimension of the query / key tensors.

  • head_dim_vo (int) – Per-head dimension of the value / output tensors.

  • page_size (int) – Page size of the paged KV-cache.

  • causal (bool) – Whether to apply a causal mask. Defaults to False.

  • sm_scale (float) – Softmax scale. If None, defaults to 1/sqrt(head_dim_qk).

  • logits_soft_cap (Optional[float]) – Logits soft-cap value. None or 0 disables capping.

  • q_data_type (torch.dtype) – Dtype of the query tensor. Defaults to torch.bfloat16.

  • kv_data_type (torch.dtype) – Dtype of the key / value tensors. Defaults to torch.bfloat16.

  • use_profiler (bool) – Whether to compile the profiler-enabled variant of the kernel. Defaults to False.

run(q: Tensor, kv_cache: Tensor | Tuple[Tensor, Tensor], out: Tensor | None = None, lse: Tensor | None = None, k_scale: Tensor | None = None, v_scale: Tensor | None = None, logits_soft_cap: float = 0.0, profiler_buffer: Tensor | None = None, kv_cache_sf: Tensor | Tuple[Tensor, Tensor] | None = None) → Tuple[Tensor, Tensor]

Run the planned holistic attention kernel.

Parameters:
  • q (torch.Tensor) – Query tensor, shape [total_qo_tokens, num_qo_heads, head_dim_qk].

  • kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) – Either a single packed paged KV-cache tensor (when K and V share storage) or a (k_cache, v_cache) pair. Layout must match the kv_layout passed to __init__().

  • out (Optional[torch.Tensor]) – Optional output buffer. If None, a new tensor is allocated with the same shape as q.

  • lse (Optional[torch.Tensor]) – Optional log-sum-exp buffer, shape [total_qo_tokens, num_qo_heads], dtype float32. Allocated if None.

  • k_scale (Optional[torch.Tensor]) – FP8 dequantization scale for k. Pre-multiplied into sm_scale.

  • v_scale (Optional[torch.Tensor]) – FP8 dequantization scale for v. Applied to the output.

  • logits_soft_cap (float) – Logits soft-cap value. Must be consistent with the logits_soft_cap passed to plan() (a non-zero value here requires a non-zero plan-time value too).

  • profiler_buffer (Optional[torch.Tensor]) – Profiler buffer. Required if the wrapper was planned with use_profiler=True.

  • kv_cache_sf (Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]]) – Optional scale tensors for NVFP4 KV-cache (one tensor or a (k_sf, v_sf) pair, mirroring the structure of kv_cache).

Returns:

(out, lse) — the attention output and its log-sum-exp.

Return type:

Tuple[torch.Tensor, torch.Tensor]

class flashinfer.attention.BatchAttentionWithAttentionSinkWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, q_data_type: dtype = torch.bfloat16, kv_data_type: dtype = torch.bfloat16, head_dim_qk: int = 128, head_dim_vo: int = 128, window_left: int = -1)

Wrapper for prefill and decode attention with paged KV-cache that adds support for attention sinks. This class extends BatchPrefillWithPagedKVCacheWrapper, providing a convenient interface for using attention sinks during prefill or decode attention.

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, q_data_type: dtype = torch.bfloat16, kv_data_type: dtype = torch.bfloat16, head_dim_qk: int = 128, head_dim_vo: int = 128, window_left: int = -1) → None

Constructor of BatchPrefillWithPagedKVCacheWrapper.

Parameters:
  • float_workspace_buffer (torch.Tensor) – The user reserved workspace buffer used to store intermediate attention results in split-k algorithm. The recommended size is 128MB, the device of the workspace buffer should be the same as the device of the input tensors. The buffer must be 16-byte aligned; tensors created by torch.empty satisfy this on supported devices.

  • kv_layout (str) – The layout of the input k/v tensors, could be either NHD or HND.

  • use_cuda_graph (bool) – Whether to enable CUDA graph capture for the prefill kernels, if enabled, the auxiliary data structures will be stored in provided buffers. The batch_size cannot change during the lifecycle of this wrapper when CUDAGraph is enabled.

  • qo_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the qo_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • paged_kv_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_indptr array, the size of this buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True.

  • paged_kv_indices_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_indices array, should be large enough to store the maximum possible size of the paged_kv_indices array during the lifetime of the wrapper. This argument is only effective when use_cuda_graph is True.

  • paged_kv_last_page_len_buf (Optional[torch.Tensor]) – The user reserved buffer to store the paged_kv_last_page_len array, the size of the buffer should be [batch_size]. This argument is only effective when use_cuda_graph is True.

  • custom_mask_buf (Optional[torch.Tensor]) – The user reserved buffer to store the custom mask tensor, should be large enough to store the maximum possible size of the packed custom mask tensor during the lifetime of the wrapper. This argument is only effective when use_cuda_graph is set to True and the custom mask will be used in attention computation.

  • mask_indptr_buf (Optional[torch.Tensor]) – The user reserved buffer to store the mask_indptr array, the size of the buffer should be [batch_size + 1]. This argument is only effective when use_cuda_graph is True and the custom mask will be used in attention computation.

  • backend (str) – The implementation backend, could be auto/fa2/fa3/cudnn/trtllm-gen cute-dsl/cute-dsl-prims. Defaults to auto. If set to auto, the wrapper will automatically choose the backend based on the device architecture and kernel availability. The cute-dsl backend uses the CuTe DSL attention kernel for Blackwell (SM100+). cute-dsl-prims backend is a SM120-only FP8 attention backend implemented using CUTLASS primitives.

  • jit_args (Optional[List[Any]]) – If provided, the wrapper will use the provided arguments to create the JIT module, otherwise, the wrapper will use default attention implementation.

  • jit_kwargs (Optional[Dict[str, Any]]) – The keyword arguments to create the JIT module, defaults to None.

  • variant_owns_mask (bool) – If True, the attention variant supplied through jit_args computes the complete attention mask in its LogitsMask hook, and MaskMode.CUSTOM is selected without a mask tensor: the kernel evaluates LogitsMask on every KV tile instead of only on the causal/window boundary tiles, and no custom_mask/packed_custom_mask needs to be passed to plan(). Requires jit_args, because the default attention variant dereferences the custom mask buffer under MaskMode.CUSTOM. Only supported with backend="fa2" (the SM90 batch prefill kernels reject MaskMode.CUSTOM), and incompatible with prefix_len_ptr (multi-item scoring), which selects a different mask mode. Defaults to False.

SM120 NVFP4 Attention

nvfp4_attention_sm120_quantize_qkv(q, k, v)

Preprocess and quantize dense Q/K/V tensors for SM120 NVFP4 attention.

nvfp4_attention_sm120_fwd(q_fp4, k_fp4, ...)

Run SM120 NVFP4 attention on pre-quantized Q/K/V tensors.

flashinfer.mla

MLA (Multi-head Latent Attention) is an attention mechanism proposed in DeepSeek series of models ( DeepSeek-V2, DeepSeek-V3, and DeepSeek-R1).

See the Batch MLA backend architecture for the planned wrapper’s ownership, metadata, tensor-representation, transactionality, and hot-path contracts.

PageAttention for MLA

trtllm_batch_decode_with_kv_cache_mla(query, ...)

Decode MLA with TRTLLM-GEN, CuteDSL, XQA, or SM120/SM121 sparse kernels.

trtllm_batch_decode_sparse_mla_dsv4(query, ...)

Decode DeepSeek V4 sparse MLA.

nvfp4_quantize_pack_sparse_mla_cache(...[, ...])

Quantize complete DeepSeek-V4 latent-KV pages to the NVFP4 cache ABI.

nvfp4_quantize_append_sparse_mla_cache(...)

Quantize and append DeepSeek-V4 latent KV by physical cache slot.

convert_compressed_page_aligned_sparse_indices_to_hca_metadata(...)

Convert combined sparse indices into reusable HCA metadata.

DSV4HCAMetadata(hca_swa_indices, ...[, ...])

Reusable metadata for the CuTe DSL causal HCA backend.

xqa_batch_decode_with_kv_cache_mla(query, ...)

XQA-backend batched MLA decode.

supported_sparse_mla_sm120_configs(*[, ...])

Enumerate the instantiated SM120 sparse-MLA decode kernel configurations.

SparseMLASm120DecodeConfig(d_qk, ...[, ...])

Instantiated decode-kernel set for one SM120 sparse-MLA kernel family.

SparseMLASm120Wrapper

alias of _SparseMLAPagedAttentionRunner

Note

With backend="cute-dsl", pass hca_swa_indices as absolute rows into the flattened SWA cache and hca_compressed_block_tables as physical compressed-cache page IDs. The SWA table has shape [B * Q, 128] and may express ring rotation or wraparound. Combined tables whose compressed segment is a canonical page expansion can opt into compatibility conversion with hca_sparse_indices_format="compressed-page-aligned". SWA entries remain arbitrary absolute rows. Precompute that conversion before a CUDA Graph or a latency-sensitive loop.

class flashinfer.mla.BatchMLAPagedAttentionWrapper(float_workspace_buffer: Tensor, use_cuda_graph: bool = False, qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, kv_indices: Tensor | None = None, kv_len_arr: Tensor | None = None, backend: str = 'auto')

Wrapper for MLA PagedAttention on DeepSeek models.

This wrapper is intended for decode and incremental prefill with the Matrix Absorption formulation of MLA, where the query/key and value/output projections are absorbed before attention. For the non-absorbed MLA prefill path, use the appropriate prefill wrapper instead.

The planned-wrapper surface owns FA2, FA3, CUTLASS, and cuTile planning and execution. Call plan() once with canonical metadata before invoking run(); the plan captures the supported input/output contract and the concrete backend’s metadata representation.

See MLA Page Layout for the paged KV-cache layout and the FlashInfer MLA blog post for the computation and Matrix Absorption background.

__init__(float_workspace_buffer: Tensor, use_cuda_graph: bool = False, qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, kv_indices: Tensor | None = None, kv_len_arr: Tensor | None = None, backend: str = 'auto') → None

Construct a planned MLA wrapper.

Parameters:
  • float_workspace_buffer (torch.Tensor) – Caller-owned workspace for intermediate attention results. A 128 MiB buffer is the usual starting point; it must be on the same device as the query and KV-cache tensors.

  • use_cuda_graph (bool, optional) – Enable CUDA-graph-compatible planning. When enabled, the optional metadata buffers below are copied at plan() time to preserve capture-time pointers. The captured batch shape cannot change.

  • qo_indptr (Optional[torch.Tensor]) – Caller-reserved int32 buffers of shape [batch_size + 1] for CSR metadata. Used only with use_cuda_graph=True.

  • kv_indptr (Optional[torch.Tensor]) – Caller-reserved int32 buffers of shape [batch_size + 1] for CSR metadata. Used only with use_cuda_graph=True.

  • kv_indices (Optional[torch.Tensor]) – Caller-reserved int32 CSR page-index buffer, sized for the maximum planned number of pages. Used only with CUDA graphs.

  • kv_len_arr (Optional[torch.Tensor]) – Caller-reserved int32 buffer of shape [batch_size] for CSR KV lengths. Used only with CUDA graphs.

  • backend ({"auto", "fa2", "fa3", "cutlass", "cutile"}) – Requested concrete backend. "auto" selects the architecture default exposed by flashinfer.utils.determine_mla_backend(). Explicit CUTLASS callers should plan with canonical dense metadata; its historical planless run path remains deprecated. Explicit cuTile callers should plan packed or split FP16/BF16 DeepSeek MLA decode inputs with canonical dense or CSR metadata. cuTile is not selected automatically.

plan(*, metadata: MLAPlanMetadata, num_heads: int, head_dim_ckv: int, head_dim_kpe: int, page_size: int, causal: bool, sm_scale: float, q_data_type: dtype, kv_data_type: dtype, use_profiler: bool = False, query_layout: Literal['packed', 'split'] = 'packed', kv_cache_layout: Literal['packed', 'split'] = 'packed', lse_mode: Literal['none', 'base2', 'basee'] = 'none', output_dtype: dtype | None = None, output_scale: Literal['none', 'per-tensor'] = 'none', scale_mode: Literal['default', 'kv-per-tensor'] = 'default', skip_softmax: bool = False) → None
plan(qo_indptr: Tensor, kv_indptr: Tensor, kv_indices: Tensor, kv_len_arr: Tensor, num_heads: int, head_dim_ckv: int, head_dim_kpe: int, page_size: int, causal: bool, sm_scale: float, q_data_type: dtype, kv_data_type: dtype, use_profiler: bool = False, *, query_layout: Literal['packed', 'split'] = 'split', kv_cache_layout: Literal['packed', 'split'] = 'split', lse_mode: Literal['none', 'base2', 'basee'] = 'none', output_dtype: dtype | None = None, output_scale: Literal['none', 'per-tensor'] = 'none', scale_mode: Literal['default', 'kv-per-tensor'] = 'default', skip_softmax: bool = False) → None
plan(*, cum_seq_lens_q: Tensor, block_tables: Tensor, seq_lens: Tensor, num_heads: int, head_dim_ckv: int, head_dim_kpe: int, page_size: int, causal: bool, sm_scale: float, q_data_type: dtype, kv_data_type: dtype, max_q_len: int | None = None, use_profiler: bool = False, query_layout: Literal['packed', 'split'] = 'split', kv_cache_layout: Literal['packed', 'split'] = 'split', lse_mode: Literal['none', 'base2', 'basee'] = 'none', output_dtype: dtype | None = None, output_scale: Literal['none', 'per-tensor'] = 'none', scale_mode: Literal['default', 'kv-per-tensor'] = 'default', skip_softmax: bool = False) → None

Plan a concrete MLA backend from canonical metadata.

Prefer metadata=MLAPlanMetadata.csr(...) for the CSR form or metadata=MLAPlanMetadata.dense(...) for the dense page-table form. MLAPlanMetadata.dual(...) may be used when both representations already exist; the planner verifies that they describe the same requests and page mapping before publishing the plan. FA2 and FA3 consume CSR metadata natively, while CUTLASS and cuTile consume dense metadata.

Metadata tensors may be on CPU or the wrapper device. They are normalized to the device required by the selected backend; tensors on another accelerator device are rejected. Passing flat CSR or dense metadata fields remains supported for compatibility, but is deprecated in favor of the metadata= object form.

The plan also declares the later run() contract. In particular, query_layout, kv_cache_layout, lse_mode, output_dtype, output_scale, and scale_mode must agree with the subsequent call. Canonical metadata defaults to packed inputs; the legacy flat forms retain their historical split-input defaults. Deprecated flat CSR plans also temporarily retain dynamic LSE behavior on FA2/FA3 and emit a warning when it is used.

Parameters:
  • metadata (Optional[MLAPlanMetadata]) – Preferred canonical CSR, dense, or dual metadata representation.

  • qo_indptr (Optional[torch.Tensor]) – Deprecated flat CSR metadata fields.

  • kv_indptr (Optional[torch.Tensor]) – Deprecated flat CSR metadata fields.

  • kv_indices (Optional[torch.Tensor]) – Deprecated flat CSR metadata fields.

  • kv_len_arr (Optional[torch.Tensor]) – Deprecated flat CSR metadata fields.

  • cum_seq_lens_q (Optional[torch.Tensor]) – Deprecated flat dense page-table metadata fields.

  • block_tables (Optional[torch.Tensor]) – Deprecated flat dense page-table metadata fields.

  • seq_lens (Optional[torch.Tensor]) – Deprecated flat dense page-table metadata fields.

  • max_q_len (Optional[int]) – Maximum dense query length; inferred from query metadata when omitted.

  • num_heads (Optional[int]) – Number of query heads.

  • head_dim_ckv (Optional[int]) – Compressed-KV and RoPE feature widths.

  • head_dim_kpe (Optional[int]) – Compressed-KV and RoPE feature widths.

  • page_size (Optional[int]) – Number of KV tokens in a cache page.

  • causal (Optional[bool]) – Whether the planned attention is causal.

  • sm_scale (Optional[float]) – Softmax scale captured by the plan.

  • q_data_type (Optional[torch.dtype]) – Query and KV-cache dtypes accepted by the selected backend.

  • kv_data_type (Optional[torch.dtype]) – Query and KV-cache dtypes accepted by the selected backend.

  • use_profiler (bool) – Whether to enable backend profiler support.

  • query_layout (Optional[{"packed", "split"}]) – Tensor representations accepted by subsequent run() calls.

  • kv_cache_layout (Optional[{"packed", "split"}]) – Tensor representations accepted by subsequent run() calls.

  • lse_mode ({"none", "base2", "basee"}) – Required log-sum-exp output mode.

  • output_dtype (Optional[torch.dtype]) – Required output dtype; defaults to q_data_type.

  • output_scale ({"none", "per-tensor"}) – Required output scaling mode.

  • scale_mode ({"default", "kv-per-tensor"}) – Required KV-scale mode for subsequent run() calls.

  • skip_softmax (bool) – Whether the plan must support the skip-softmax threshold feature.

Notes

Positional arguments are deprecated; use keyword arguments for all plan() parameters. Flat metadata arguments are also deprecated; use an MLAPlanMetadata object instead.

run(*, query: object, kv_cache: object, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False, o_scale: float | None = None, ckv_scale: float | None = None, ckv_scale_arr: Tensor | None = None, kpe_scale: float | None = None) → Tensor
run(*, query: object, kv_cache: object, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True], profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False, o_scale: float | None = None, ckv_scale: float | None = None, ckv_scale_arr: Tensor | None = None, kpe_scale: float | None = None) → Tuple[Tensor, Tensor]
run(q_nope: Tensor | None = None, q_pe: Tensor | None = None, ckv_cache: Tensor | None = None, kpe_cache: Tensor | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False, o_scale: float | None = None, *, ckv_scale: float | None = None, ckv_scale_arr: Tensor | None = None, kpe_scale: float | None = None) → Tensor
run(q_nope: Tensor | None = None, q_pe: Tensor | None = None, ckv_cache: Tensor | None = None, kpe_cache: Tensor | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False, o_scale: float | None = None, *, ckv_scale: float | None = None, ckv_scale_arr: Tensor | None = None, kpe_scale: float | None = None) → Tuple[Tensor, Tensor]

Run one planned MLA attention request.

Preferred structural input form

Pass query= and kv_cache=. Each may be a packed tensor whose final dimension joins its no-PE and PE features, or a split pair of tensors. The representation must satisfy the layouts declared by plan(). Adjacent split views can be reinterpreted as packed tensors without a copy; FA2 and FA3 also accept independent split inputs natively. Packed-native CUTLASS plans reject independent split KV caches instead of silently copying them.

Output forms

With return_lse=False (the default), the method returns the output tensor. With return_lse=True, it returns (output, lse) when the canonical plan was created with a matching LSE mode. return_lse_base_on_e selects natural-log rather than base-2 LSE output. Deprecated flat CSR FA2/FA3 plans temporarily retain dynamic LSE behavior and emit a warning. Caller-provided out and lse buffers are used directly and returned by identity.

Legacy separate-parameter form

The separate q_nope / q_pe and ckv_cache / kpe_cache parameters remain a deprecated compatibility form. Pass the same split tensors through query=(q_nope, q_pe) and kv_cache=(ckv_cache, kpe_cache) instead. Structural split values are not deprecated. Positional arguments are deprecated independently; each warning is emitted once per wrapper instance.

Parameters:
  • query (object) – Preferred packed or split structural query value.

  • kv_cache (object) – Preferred packed or split structural KV-cache value.

  • q_nope (Optional[torch.Tensor]) – Legacy split query tensors. Supply both together.

  • q_pe (Optional[torch.Tensor]) – Legacy split query tensors. Supply both together.

  • ckv_cache (Optional[torch.Tensor]) – Legacy split compressed-KV and PE cache tensors. Supply both together.

  • kpe_cache (Optional[torch.Tensor]) – Legacy split compressed-KV and PE cache tensors. Supply both together.

  • out (Optional[torch.Tensor]) – Caller-owned output buffer. When o_scale is provided, it must be an FP8 tensor for CUTLASS output.

  • lse (Optional[torch.Tensor]) – Caller-owned log-sum-exp output buffer for an LSE-enabled plan.

  • return_lse (bool) – Return the LSE tensor in addition to the output.

  • profiler_buffer (Optional[torch.Tensor]) – Backend profiler output buffer.

  • kv_len (Optional[torch.Tensor]) – CUTLASS/cuTile metadata aliases. A planned request may omit them; the deprecated unplanned CUTLASS path requires both. Runtime metadata is a trusted hot-path input: every length must be nonnegative and fit within its page-table row, every live page ID must index kv_cache, and callers must not mutate either tensor while a launch is in flight. These values are not synchronized to the host for validation so that CUDA-graph capture remains valid.

  • page_table (Optional[torch.Tensor]) – CUTLASS/cuTile metadata aliases. A planned request may omit them; the deprecated unplanned CUTLASS path requires both. Runtime metadata is a trusted hot-path input: every length must be nonnegative and fit within its page-table row, every live page ID must index kv_cache, and callers must not mutate either tensor while a launch is in flight. These values are not synchronized to the host for validation so that CUDA-graph capture remains valid.

  • return_lse_base_on_e (bool) – Return natural-log rather than base-2 LSE values.

  • o_scale (Optional[float]) – Per-tensor FP8 output scale supported by CUTLASS.

  • ckv_scale (optional) – Per-tensor or per-token FP8 KV-cache scales required by plans that selected scale_mode="kv-per-tensor".

  • ckv_scale_arr (optional) – Per-tensor or per-token FP8 KV-cache scales required by plans that selected scale_mode="kv-per-tensor".

  • kpe_scale (optional) – Per-tensor or per-token FP8 KV-cache scales required by plans that selected scale_mode="kv-per-tensor".

Notes

Positional arguments are deprecated; use keyword arguments for plan() and run(). The separate q_nope / q_pe and ckv_cache / kpe_cache parameters are also deprecated; pass structural query= and kv_cache= values instead. An explicitly requested CUTLASS backend may still run without plan() when both kv_len and page_table are supplied, but that compatibility path is also deprecated.

class flashinfer.mla.MLAPlanMetadata(qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, kv_indices: Tensor | None = None, kv_len_arr: Tensor | None = None, cum_seq_lens_q: Tensor | None = None, block_tables: Tensor | None = None, seq_lens: Tensor | None = None, max_q_len: int | None = None)

Canonical CSR and/or dense metadata for a Batch MLA plan.

BatchMLAPagedAttentionWrapper.plan accepts the preferred metadata=MLAPlanMetadata.csr(...) or metadata=MLAPlanMetadata.dense(...) keyword form. The keyword form defaults run inputs to packed query / kv_cache structural tensors. The legacy flat CSR and dense metadata adapters remain available for compatibility, emit a DeprecationWarning once per process, and default run inputs to split q_nope / q_pe and ckv_cache / kpe_cache. Structural run inputs use an exact tuple grammar: a tensor is packed, (left, right) is split, and (packed, (left, right)) or ((left, right), packed) is a trusted redundant form. The selected form must match the planned rank, leading shape, dtype, device, and split widths.

LSE mode, output dtype/scaling, and KV scaling are plan/run contracts. A run that needs different values must re-plan first, except that deprecated flat CSR FA2/FA3 plans temporarily preserve dynamic LSE behavior with a DeprecationWarning. Explicit backend="cutlass" callers that omit plan remain supported through a deprecated compatibility adapter when kv_len and page_table are supplied.