flashinfer.cudnn.cudnn_batch_decode_with_kv_cache¶
- flashinfer.cudnn.cudnn_batch_decode_with_kv_cache(q: Tensor, k_cache: Tensor, v_cache: Tensor, scale: float, workspace_buffer: Tensor, *, max_sequence_kv: int, actual_seq_lens_kv: Tensor | None = None, block_tables: Tensor | None = None, is_cuda_graph_compatible: bool = False, batch_offsets_q: Tensor | None = None, batch_offsets_o: Tensor | None = None, batch_offsets_k: Tensor | None = None, batch_offsets_v: Tensor | None = None, out: Tensor | None = None, return_lse: bool = False, lse: Tensor | None = None, q_len_per_req: int = 1, window_left: int = -1, sinks: Tensor | None = None) Tensor | tuple[Tensor, Tensor]¶
Batched decode attention with paged KV cache, backed by cuDNN SDPA.
- Parameters:
q (torch.Tensor) – Query tensor of shape
(batch_size * q_len_per_req, num_heads_qo, head_dim)((batch_size, num_heads_qo, head_dim)for plain one-token decode),torch.float16ortorch.bfloat16(the output usesq.dtype). Withq_len_per_req > 1the rows of one request are consecutive.torch.float16requires the cuDNN graph backend; the fallback (cubin) path is bf16-only and raisesNotImplementedError.k_cache (torch.Tensor) – Key cache, shape
(total_num_pages, num_heads_kv, page_size, head_dim).v_cache (torch.Tensor) – Value cache, shape
(total_num_pages, num_heads_kv, page_size, head_dim).scale (float) – Softmax scaling factor, typically
1 / sqrt(head_dim).workspace_buffer (torch.Tensor) – Workspace buffer for cuDNN. Scales with batch size; 128 MB is sufficient for typical decode workloads.
max_sequence_kv (int) – Maximum number of tokens per KV sequence in the batch (
s_kv_max).actual_seq_lens_kv (Optional[torch.Tensor]) – Per-request KV lengths, shape
(batch_size,). When cuDNN is available (the default backend) this tensor must reside on the same CUDA device asq. Only the fallback non-cuDNN path accepts (and internally copies) a CPU tensor.block_tables (Optional[torch.Tensor]) – Page-table mapping for the paged KV cache, shape
(batch_size, num_pages_per_seq)on GPU.is_cuda_graph_compatible (bool) – Whether to plan the operation in a CUDA-graph-capture-safe mode.
batch_offsets_q (Optional[torch.Tensor]) – Per-request element offsets into the query tensor, int32 or int64, shape
(batch_size,)or(batch_size + 1,)on GPU (optional end offset). The cuDNN graph path accepts only dense offsets matching the query’s batch stride. PreferNoneto avoid redundant device-side checks.batch_offsets_o (Optional[torch.Tensor]) – Like
batch_offsets_q, but for the contiguous output tensor. On the cuDNN graph path, non-dense Q/O offsets trigger an asynchronous device assertion, including if changed before CUDA graph replay.batch_offsets_k (Optional[torch.Tensor]) – Per-request offsets into the key tensor, shape
(batch_size,)on GPU.batch_offsets_v (Optional[torch.Tensor]) – Per-request offsets into the value tensor, shape
(batch_size,)on GPU.out (Optional[torch.Tensor]) – Pre-allocated output tensor, shape
(batch_size, num_heads_qo, head_dim)with dtypeq.dtype; allocated internally whenNone.return_lse (bool) – Whether to also return the log-sum-exp of the attention scores (cuDNN’s SDPA
Statsoutput). Requires the cuDNN graph backend; raisesNotImplementedErroron the fallback (cubin) path.lse (Optional[torch.Tensor]) – Pre-allocated LSE tensor, shape
(batch_size * q_len_per_req, num_heads_qo),torch.float32, contiguous, on the same device asq; allocated internally whenNoneandreturn_lseisTrue.q_len_per_req (int) – Query rows per request (speculative / multi-token-prediction verification). Rows of one request attend under the bottom-right causal diagonal: row
iof a request withkv_lenkeys sees keys0 .. kv_len - q_len_per_req + i. Every request needskv_len >= q_len_per_req. Defaults to1(no mask beyond padding).window_left (int) – Left sliding-window bound in FlashInfer’s convention: a row attends to the
window_leftkeys before its diagonal position plus that position itself;-1(default) disables the window.sinks (Optional[torch.Tensor]) – Per-head attention sink logits, shape
(num_heads_qo,),torch.float32, onq’s device.sinks[h]joins each row’s softmax denominator as one extra logit with a zero value row, as in FlashInfer’s other backends (gpt-oss / Streaming-LLM sinks). Whether the cuDNN stack serves a sink atq_len_per_req == 1is decided by its SDPA engines (cudnn-frontend 1.30+ with the FROST engines enabled does; the backend engine raises a not-supported error at graph build).
- Returns:
Output tensor of shape
(batch_size * q_len_per_req, num_heads_qo, head_dim)whenreturn_lse=False; otherwise(output, lse)wherelsehas shape(batch_size * q_len_per_req, num_heads_qo)and dtypetorch.float32.- Return type:
Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]
Note
All tensors must be on the same CUDA device.
qmay carry arbitrary batch/head strides (e.g. a slice of a packed QKV projection) as long ashead_dimis innermost and dense;out/lsemust be contiguous. Query and KV heads may differ (num_heads_qo >= num_heads_kv, multi-query / grouped-query attention).LSE convention:
lse[b, h]is the base-2 log-sum-exp of the pre-softmax attention row withscalefolded in, i.e.log2(sum_j(exp(scale * q[b, h] . k[b, h // (num_heads_qo // num_heads_kv), j])))summed over the valid KV positionsj < actual_seq_lens_kv[b]— the same contract as every other FlashInfer backend (torch.logsumexp(...) * log2(e)), so it can be fed to the cascade-merge kernels. cuDNN emits natural-log stats; they are folded to base-2 here.