flashinfer.cudnn.cudnn_chunk_gated_delta_product

flashinfer.cudnn.cudnn_chunk_gated_delta_product(q: Tensor, k: Tensor, v: Tensor, g: Tensor | None = None, beta: Tensor | None = None, num_householder: int = 1, scale: float | None = None, initial_state: Tensor | None = None, output_final_state: bool = False, cu_seqlens: Tensor | None = None, use_qk_l2norm_in_kernel: bool = False, output: Tensor | None = None, output_state: Tensor | None = None, batch_invariant: bool = False) → Tensor | tuple[Tensor, Tensor]

Chunked Gated DeltaProduct prefill on cuDNN’s fused SM100 engine.

GDP applies n = num_householder beta-gated Householder updates per token with one per-head scalar decay per token: the GDN recurrence on an expanded sub-token timeline, with the decay acting before the token’s updates and the readout following the last one. num_householder == 1 is exactly cudnn_chunk_gated_delta_rule().

Requires cudnn-frontend 1.29+ with the cutedsl extra. Everything else the engine decides for itself: it declines a graph it cannot serve (the per-engine reason lands in the frontend’s log).

Parameters:
  • q (torch.Tensor) – [total_seq_len, num_q_heads, head_size], packed at real-token rows. Strides are passed through to cuDNN, so only the innermost dim has to be contiguous.

  • k (torch.Tensor) – [total_seq_len * num_householder, num_k_heads / num_v_heads, head_size], packed on the expanded sub-token timeline: the n Householder updates of token t occupy rows t*n .. t*n + n - 1. num_k_heads must equal num_q_heads or num_v_heads.

  • v (torch.Tensor) – [total_seq_len * num_householder, num_k_heads / num_v_heads, head_size], packed on the expanded sub-token timeline: the n Householder updates of token t occupy rows t*n .. t*n + n - 1. num_k_heads must equal num_q_heads or num_v_heads.

  • g (torch.Tensor, optional) – Per-head forget gate in linear space (alpha = exp(log_g)), shape [total_seq_len, num_sab_heads] at real-token rows, passed to cuDNN as is (gate_domain="linear") at g’s own dtype. All-ones when None.

  • beta (torch.Tensor, optional) – Per-head, per-Householder update gate [total_seq_len * num_householder, num_sab_heads], post-sigmoid, in float32 or q.dtype. All-ones when None.

  • num_householder (int) – Householder updates per token (n >= 1).

  • scale (float, optional) – Query scale; 1 / sqrt(head_size) when None or 0.0, matching the native GDN path.

  • initial_state (torch.Tensor, optional) – State [num_seqs, num_sab_heads, head_size, head_size], V-major, float32 or bfloat16. output_state must not alias initial_state; see cudnn_chunk_gated_delta_rule().

  • output_state (torch.Tensor, optional) – State [num_seqs, num_sab_heads, head_size, head_size], V-major, float32 or bfloat16. output_state must not alias initial_state; see cudnn_chunk_gated_delta_rule().

  • output_final_state (bool) – Return the outgoing recurrent state alongside the output; when unset no state is written at all.

  • cu_seqlens (torch.Tensor) – [num_seqs + 1] int32 or int64, over the real tokens. Required.

  • use_qk_l2norm_in_kernel (bool) – Fuse the q/k L2 normalization into the kernel.

  • output (torch.Tensor, optional) – Pre-allocated [total_seq_len, num_o_heads, head_size] at real-token rows, written in place by the kernel.

  • batch_invariant (bool) – Disable the split-K partition; see cudnn_chunk_gated_delta_rule().

Returns:

output, or (output, final_state) when output_final_state.

Return type:

torch.Tensor or Tuple[torch.Tensor, torch.Tensor]