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.float16 or torch.bfloat16 (the output uses q.dtype). With q_len_per_req > 1 the rows of one request are consecutive. torch.float16 requires the cuDNN graph backend; the fallback (cubin) path is bf16-only and raises NotImplementedError.

  • 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 as q. 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. Prefer None to 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 dtype q.dtype; allocated internally when None.

  • return_lse (bool) – Whether to also return the log-sum-exp of the attention scores (cuDNN’s SDPA Stats output). Requires the cuDNN graph backend; raises NotImplementedError on 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 as q; allocated internally when None and return_lse is True.

  • 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 i of a request with kv_len keys sees keys 0 .. kv_len - q_len_per_req + i. Every request needs kv_len >= q_len_per_req. Defaults to 1 (no mask beyond padding).

  • window_left (int) – Left sliding-window bound in FlashInfer’s convention: a row attends to the window_left keys 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, on q’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 at q_len_per_req == 1 is 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) when return_lse=False; otherwise (output, lse) where lse has shape (batch_size * q_len_per_req, num_heads_qo) and dtype torch.float32.

Return type:

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

Note

All tensors must be on the same CUDA device. q may carry arbitrary batch/head strides (e.g. a slice of a packed QKV projection) as long as head_dim is innermost and dense; out/lse must 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 with scale folded 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 positions j < 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.