flashinfer.gdp_prefill.chunk_gated_delta_product

flashinfer.gdp_prefill.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, *, backend: Literal['auto', 'cudnn'] = 'auto') → Tensor | Tuple[Tensor, Tensor]

Chunked Gated DeltaProduct (GDP) attention for prefill.

GDP generalizes flashinfer.chunk_gated_delta_rule() from one beta-gated Householder update per token to n = num_householder of them, with one per-head scalar decay per token: the decay acts before the token’s updates and the readout follows the last one. Per real token, with sub-token rows t*n .. t*n + n - 1 of k / v / beta,

\[\begin{split}S &\mathrel{*}= \alpha_t \\ v^{new}_{t,j} &= \beta_{t,j} (v_{t,j} - S k_{t,j}), \quad S \mathrel{+}= v^{new}_{t,j} \otimes k_{t,j}, \quad j = 0 \ldots n - 1 \\ o_t &= \mathrm{scale} \cdot S q_t\end{split}\]

num_householder == 1 is exactly flashinfer.chunk_gated_delta_rule().

Only varlen (packed / THD) input is accepted, matching flashinfer.chunk_gated_delta_rule(): pass cu_seqlens over the real tokens and lay the tokens of all sequences end to end.

Parameters:
  • q (torch.Tensor) – [total_seq_len, num_q_heads, head_size], float16 or bfloat16, at real-token rows. Strides are honored, so only the innermost dimension has to be contiguous.

  • k (torch.Tensor) – [total_seq_len * num_householder, num_k_heads / num_v_heads, head_size], float16 or bfloat16, 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], float16 or bfloat16, 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, elementwise in (0, 1]), shape [total_seq_len, num_sab_heads] at real-token rows, at float32, bfloat16 or float16. All-ones (no decay) 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 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 real-token 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] at real-token rows, 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 GDP 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_product().

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).