flashinfer.cudnn.cudnn_chunk_gated_delta_rule¶
- flashinfer.cudnn.cudnn_chunk_gated_delta_rule(q: Tensor, k: Tensor, v: Tensor, g: Tensor | None = None, beta: 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, batch_invariant: bool = False) Tensor | tuple[Tensor, Tensor]¶
Chunked Gated Delta Rule prefill on cuDNN’s fused SM100 engine.
Argument meanings match
flashinfer.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 / num_k_heads / num_v_heads, 128], packed. Strides are passed through to cuDNN, so only the innermost dim has to be contiguous.k (torch.Tensor) –
[total_seq_len, num_q_heads / num_k_heads / num_v_heads, 128], packed. Strides are passed through to cuDNN, so only the innermost dim has to be contiguous.v (torch.Tensor) –
[total_seq_len, num_q_heads / num_k_heads / num_v_heads, 128], packed. Strides are passed through to cuDNN, so only the innermost dim has to be contiguous.g (torch.Tensor, optional) – Per-head forget gate in linear space (
alpha = exp(log_g)), shape[total_seq_len, num_sab_heads], passed to cuDNN as is (gate_domain="linear") atg’s own dtype, which cuDNN reads at float32, bfloat16 or float16. All-ones whenNone.beta (torch.Tensor, optional) – Per-head update gate
[total_seq_len, num_sab_heads], post-sigmoid, in float32 orq.dtype. All-ones whenNone.scale (float, optional) – Query scale;
1 / sqrt(head_dim)whenNoneor0.0, matching the native GDN path.initial_state (torch.Tensor, optional) – State
[num_seqs, num_sab_heads, 128, 128], V-major, float32 or bfloat16. cuDNN uses the same layout, so these pass through untransposed andoutput_stateis written in place by the kernel.output_statemust not aliasinitial_state: the engines split one sequence across CTAs, so the chunk-0 CTA reading the incoming state would race the last-chunk CTA writing the outgoing one. Like the state-slot uniquenessstate_indicesrelies on, this is a caller precondition rather than a launch-time check.output_state (torch.Tensor, optional) – State
[num_seqs, num_sab_heads, 128, 128], V-major, float32 or bfloat16. cuDNN uses the same layout, so these pass through untransposed andoutput_stateis written in place by the kernel.output_statemust not aliasinitial_state: the engines split one sequence across CTAs, so the chunk-0 CTA reading the incoming state would race the last-chunk CTA writing the outgoing one. Like the state-slot uniquenessstate_indicesrelies on, this is a caller precondition rather than a launch-time check.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. 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, 128], written in place by the kernel.batch_invariant (bool) – Disable the split-K partition so the reduction order, and hence the result, does not depend on how sequences are batched. Costs the parallelism split-K exists to create on few long sequences and saves its fixed scheduling cost on many short ones.
- Returns:
output, or(output, final_state)whenoutput_final_state.- Return type:
torch.Tensor or Tuple[torch.Tensor, torch.Tensor]