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 ton = num_householderof 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 rowst*n .. t*n + n - 1ofk/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 == 1is exactlyflashinfer.chunk_gated_delta_rule().Only varlen (packed / THD) input is accepted, matching
flashinfer.chunk_gated_delta_rule(): passcu_seqlensover 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: thenHouseholder updates of tokentoccupy rowst*n .. t*n + n - 1.num_k_headsmust equalnum_q_headsornum_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: thenHouseholder updates of tokentoccupy rowst*n .. t*n + n - 1.num_k_headsmust equalnum_q_headsornum_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) whenNone.beta (torch.Tensor, optional) – Per-head, per-Householder update gate
[total_seq_len * num_householder, num_sab_heads], post-sigmoid, in float32 orq.dtype. All-ones whenNone.num_householder (int) – Householder updates per token (
n >= 1).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 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 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 GDP kernel of its own, so
"auto"(default) and"cudnn"both run cuDNN’s fused SM100 linear-attention engine throughflashinfer.cudnn.cudnn_chunk_gated_delta_product().
- 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).