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 cutedsl extra. 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") at g’s own dtype, which cuDNN reads at float32, bfloat16 or float16. All-ones when None.

  • beta (torch.Tensor, optional) – Per-head update gate [total_seq_len, num_sab_heads], post-sigmoid, in float32 or q.dtype. All-ones when None.

  • scale (float, optional) – Query scale; 1 / sqrt(head_dim) when None or 0.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 and output_state is written in place by the kernel. output_state must not alias initial_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 uniqueness state_indices relies 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 and output_state is written in place by the kernel. output_state must not alias initial_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 uniqueness state_indices relies 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) when output_final_state.

Return type:

torch.Tensor or Tuple[torch.Tensor, torch.Tensor]