flashinfer.cudnn.cudnn_chunk_gated_delta_product¶
- flashinfer.cudnn.cudnn_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, batch_invariant: bool = False) Tensor | tuple[Tensor, Tensor]¶
Chunked Gated DeltaProduct prefill on cuDNN’s fused SM100 engine.
GDP applies
n = num_householderbeta-gated Householder updates per token with one per-head scalar decay per token: the GDN recurrence on an expanded sub-token timeline, with the decay acting before the token’s updates and the readout following the last one.num_householder == 1is exactlycudnn_chunk_gated_delta_rule().Requires cudnn-frontend 1.29+ with the
cutedslextra. Everything else the engine decides for itself: it declines a graph it cannot serve (the per-engine reason lands in the frontend’s log).- Parameters:
q (torch.Tensor) –
[total_seq_len, num_q_heads, head_size], packed at real-token rows. Strides are passed through to cuDNN, so only the innermost dim has to be contiguous.k (torch.Tensor) –
[total_seq_len * num_householder, num_k_heads / num_v_heads, head_size], packed 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], packed 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 = exp(log_g)), shape[total_seq_len, num_sab_heads]at real-token rows, passed to cuDNN as is (gate_domain="linear") atg’s own dtype. All-ones 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, matching the native GDN path.initial_state (torch.Tensor, optional) – State
[num_seqs, num_sab_heads, head_size, head_size], 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, head_size, head_size], 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, over the real tokens. 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, head_size]at real-token rows, written in place by the kernel.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]