flashinfer.cudnn.cudnn_chunk_gated_delta_rule2

flashinfer.cudnn.cudnn_chunk_gated_delta_rule2(q: Tensor, k: Tensor, v: Tensor, g: Tensor | None = None, beta: Tensor | None = None, w: Tensor | None = None, 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 Delta Rule 2 prefill on cuDNN’s fused SM100 engine.

Argument meanings match flashinfer.chunk_gated_delta_rule2().

GDN-2 generalizes GDN’s per-head scalar gates to channel-wise ones: the forget gate g and erase gate beta are per key channel and the write gate w is per value channel.

\[\begin{split}S_t &= \mathrm{diag}(e^{g_t}) S_{t-1} \\ v^{new}_t &= w_t \odot v_t - (\beta_t \odot k_t)^\top S_t \\ S_t &\mathrel{+}= k_t \otimes v^{new}_t \\ o_t &= \mathrm{scale} \cdot q_t^\top S_t\end{split}\]

Requires cudnn-frontend 1.29+ with the cutedsl extra. Everything else the engine decides for itself.

Parameters:
  • q (torch.Tensor) – [total_seq_len, num_q_heads / num_k_heads / num_v_heads, 128], packed.

  • k (torch.Tensor) – [total_seq_len, num_q_heads / num_k_heads / num_v_heads, 128], packed.

  • v (torch.Tensor) – [total_seq_len, num_q_heads / num_k_heads / num_v_heads, 128], packed.

  • g (torch.Tensor, optional) – Channel-wise forget gate as the natural-log decay, shape [total_seq_len, num_sab_heads, 128], passed to cuDNN as is (gate_domain="log") at g’s own dtype, which cuDNN reads at float32, bfloat16 or float16. All-zeros when None.

  • beta (torch.Tensor, optional) – Channel-wise erase gate [total_seq_len, num_sab_heads, 128], converted to q.dtype. All-ones when None.

  • w (torch.Tensor, optional) – Channel-wise write gate [total_seq_len, num_sab_heads, 128], converted to q.dtype. All-ones when None.

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

  • initial_state (torch.Tensor, optional) – State [num_seqs, num_sab_heads, 128, 128], 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, 128, 128], 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. 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, 128].

  • 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]