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
gand erase gatebetaare per key channel and the write gatewis 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
cutedslextra. 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") atg’s own dtype, which cuDNN reads at float32, bfloat16 or float16. All-zeros whenNone.beta (torch.Tensor, optional) – Channel-wise erase gate
[total_seq_len, num_sab_heads, 128], converted toq.dtype. All-ones whenNone.w (torch.Tensor, optional) – Channel-wise write gate
[total_seq_len, num_sab_heads, 128], converted toq.dtype. All-ones whenNone.scale (float, optional) – Query scale;
1 / sqrt(head_dim)whenNoneor0.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_statemust not aliasinitial_state; seecudnn_chunk_gated_delta_rule().output_state (torch.Tensor, optional) – State
[num_seqs, num_sab_heads, 128, 128], V-major, float32 or bfloat16.output_statemust not aliasinitial_state; seecudnn_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)whenoutput_final_state.- Return type:
torch.Tensor or Tuple[torch.Tensor, torch.Tensor]