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 gate g and the erase gate beta are per key channel, and a third gate w scales 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, beta and w to 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(): pass cu_seqlens and 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_heads must equal num_q_heads or num_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_heads must equal num_q_heads or num_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_heads must equal num_q_heads or num_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) when None.

  • beta (torch.Tensor, optional) – Channel-wise erase gate [total_seq_len, num_sab_heads, head_size], post-sigmoid, read at q.dtype. All-ones when None.

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

  • scale (float, optional) – Query scale; 1 / sqrt(head_size) when None or 0.0, matching flashinfer.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 when None.

  • 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 when None.

  • output_state (torch.Tensor, optional) – Pre-allocated final-state buffer, written in place. Allocated internally when None and output_final_state is set. It must not alias initial_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 through flashinfer.cudnn.cudnn_chunk_gated_delta_rule2().

Returns:

output, or (output, final_state) when output_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 cutedsl extra (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).