flashinfer.msa_ops.prepare_msa_nvfp4_sparse_decode

flashinfer.msa_ops.prepare_msa_nvfp4_sparse_decode(q: Tensor, k: Tensor, v: Tensor, q2k_indices: Tensor, *, k_scale: Tensor, v_scale: Tensor, page_table: Tensor, seqused_k: Tensor, k_global_scale: float, v_global_scale: float, workspace_buffer: Tensor | None = None, seqlen_q: int = 1, softmax_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, backend: str = 'cake')

Warning

NVFP4 paged-KV MSA decode (Cake backend) is experimental: it provides no compatibility guarantees and may change or be removed without deprecation.

Prepare NVFP4 paged-KV sparse decode with the generated Cake program (SM100/SM103).

The experimental Cake backend serves the same problem surface as the packed-NVFP4 paged-KV route of msa_sparse_decode_attention() – the planar page pool of docs/design_docs/nvfp4_msa_paged_kv_layout.md, top-k 16, right-aligned causal decode tokens – with one persistent kernel that keeps the K tiles resident in tensor memory and runs both MMAs in the swapped orientation. Preparation validates and binds the tensors; the returned runner launches with no allocation and no host synchronization and can be captured into a CUDA Graph.

qtorch.Tensor

BF16 [batch * seqlen_q, num_q_heads, 128]; request b owns rows [b * seqlen_q, (b + 1) * seqlen_q) and query head h attends to KV head h // (num_q_heads // num_kv_heads) (at most sixteen query heads per KV head).

k, vtorch.Tensor

uint8 [num_pages, num_kv_heads, 128, 64] strided views of the planar page pool (packed E2M1, two values per byte).

q2k_indicestorch.Tensor

int32 [num_kv_heads, batch * seqlen_q, 16] selected page indices per KV head and query token, ascending and -1 padded, as msa_topk_select() produces them.

k_scale, v_scaletorch.Tensor

[num_pages, num_kv_heads, 128, 8] strided views of the E4M3 block scales (uint8 or float8_e4m3fn); K scales linear, V scales in the cache writer’s swizzled order.

page_tabletorch.Tensor

int32 [batch, max_pages] physical page ids per request.

seqused_ktorch.Tensor

int32 [batch] KV tokens per request including the new tokens.

k_global_scale, v_global_scalefloat

Positive per-side global scales; the K scale is folded into the softmax scale and the V scale is applied in the epilogue.

workspace_bufferOptional[torch.Tensor]

Caller-owned CUDA bytes for the split-KV partials, at least flashinfer.experimental.msa_nvfp4_decode.cake_backend.msa_nvfp4_decode_workspace_size(); required only when the batch is small enough to split (the runner reports the chosen factor as splits). The kernel resets its completion counters after every merge, so the region must not be shared with other work between launches.

seqlen_qint

Query tokens per request, in [1, 32]; token i sits at KV position seqused_k[b] - seqlen_q + i and attends causally.

softmax_scaleOptional[float]

Defaults to 1 / sqrt(128).

out, lseOptional[torch.Tensor]

Optional caller-owned BF16 output [batch * seqlen_q, num_q_heads, 128] and float32 natural-log softmax normalizer [batch * seqlen_q, num_q_heads].

backendstr

Only "cake" is supported.

MSANvfp4DecodeRunner

Calling it launches the decode on the current stream and returns out. See flashinfer/experimental/msa_nvfp4_decode/README.md.