flashinfer.gdn2_prefill.chunk_gated_delta_rule2¶
- flashinfer.gdn2_prefill.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, *, backend: Literal['auto', 'cudnn'] = 'auto') Tensor | Tuple[Tensor, Tensor]¶
Chunked Gated Delta Rule 2 (GDN-2) attention for prefill.
GDN-2 generalizes
flashinfer.chunk_gated_delta_rule()’s per-head scalar gates to channel-wise ones: the forget gategand the erase gatebetaare per key channel, and a third gatewscales the incoming value 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}\]Setting
g,betaandwto per-channel constants recovers GDN: \(w \equiv \beta\) and a channel-constant \(g\) make the update the scalar-gated delta rule.Only varlen (packed / THD) input is accepted, matching
flashinfer.chunk_gated_delta_rule(): passcu_seqlensand lay the tokens of all sequences end to end.- Parameters:
q (torch.Tensor) –
[total_seq_len, num_q_heads / num_k_heads / num_v_heads, head_size], float16 or bfloat16.num_k_headsmust equalnum_q_headsornum_v_heads, and the query and value head counts must be equal or one a multiple of the other. Strides are honored, so only the innermost dimension has to be contiguous.k (torch.Tensor) –
[total_seq_len, num_q_heads / num_k_heads / num_v_heads, head_size], float16 or bfloat16.num_k_headsmust equalnum_q_headsornum_v_heads, and the query and value head counts must be equal or one a multiple of the other. Strides are honored, so only the innermost dimension has to be contiguous.v (torch.Tensor) –
[total_seq_len, num_q_heads / num_k_heads / num_v_heads, head_size], float16 or bfloat16.num_k_headsmust equalnum_q_headsornum_v_heads, and the query and value head counts must be equal or one a multiple of the other. Strides are honored, so only the innermost dimension has to be contiguous.g (torch.Tensor, optional) – Channel-wise forget gate as the natural-log decay (
ln alpha, elementwise<= 0), shape[total_seq_len, num_sab_heads, head_size], at float32, bfloat16 or float16. All-zeros (no decay) whenNone.beta (torch.Tensor, optional) – Channel-wise erase gate
[total_seq_len, num_sab_heads, head_size], post-sigmoid, read atq.dtype. All-ones whenNone.w (torch.Tensor, optional) – Channel-wise write gate
[total_seq_len, num_sab_heads, head_size], read atq.dtype. All-ones whenNone.scale (float, optional) – Query scale;
1 / sqrt(head_size)whenNoneor0.0, matchingflashinfer.chunk_gated_delta_rule().initial_state (torch.Tensor, optional) – Incoming recurrent state
[num_seqs, num_sab_heads, head_size, head_size], V-major ([..., V, K]), float32 or bfloat16. A zero state whenNone.output_final_state (bool) – Return the outgoing recurrent state alongside the output.
cu_seqlens (torch.Tensor) – Cumulative sequence lengths
[num_seqs + 1], int32 or int64. Required.use_qk_l2norm_in_kernel (bool) – Fuse the q/k L2 normalization into the kernel instead of expecting pre-normalized inputs.
output (torch.Tensor, optional) – Pre-allocated output
[total_seq_len, num_o_heads, head_size], written in place. Allocated internally whenNone.output_state (torch.Tensor, optional) – Pre-allocated final-state buffer, written in place. Allocated internally when
Noneandoutput_final_stateis set. It must not aliasinitial_state; the kernel splits one sequence across CTAs, so the CTA reading the incoming state would race the one writing the outgoing state.backend (Literal["auto", "cudnn"], optional) – FlashInfer carries no GDN-2 kernel of its own, so
"auto"(default) and"cudnn"both run cuDNN’s fused SM100 linear-attention engine throughflashinfer.cudnn.cudnn_chunk_gated_delta_rule2().
- Returns:
output, or(output, final_state)whenoutput_final_state.- Return type:
torch.Tensor or Tuple[torch.Tensor, torch.Tensor]
Note
Requires an SM100-family (Blackwell) device and cudnn-frontend 1.29+ with the
cutedslextra (pip install 'nvidia-cudnn-frontend[cutedsl]'). Everything finer – head dims, input dtypes, head-count relations – is the engine’s call: a graph it cannot serve is declined by cuDNN (the per-engine reason lands in the frontend’s log).