flashinfer.cudnn.cudnn_recurrent_kda

flashinfer.cudnn.cudnn_recurrent_kda(q: Tensor, k: Tensor, v: Tensor, g: Tensor, beta: Tensor, A_log: Tensor | None = None, dt_bias: Tensor | None = None, scale: float | None = None, initial_state: Tensor | None = None, output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = True, use_gate_in_kernel: bool = False, lower_bound: float | None = None, cu_seqlens: Tensor | None = None, beta_is_logit: bool = False, output: Tensor | None = None, output_state: Tensor | None = None, batch_invariant: bool = False) → tuple[Tensor, Tensor | None]

Kimi Delta Attention prefill on cuDNN’s fused SM100 engine.

Argument meanings match flashinfer.recurrent_kda(), restricted to the ordinary multi-token prefill subset: no speculative decode, no state pool, no initial_state_source, no state checkpoints.

Requires cudnn-frontend 1.29+ with the cutedsl extra. Everything else the engine decides for itself.

Parameters:
  • q (torch.Tensor) – [1, total_tokens, H, 128] or [total_tokens, H, 128], bfloat16 or float16.

  • k (torch.Tensor) – [1, total_tokens, H, 128] or [total_tokens, H, 128], bfloat16 or float16.

  • v (torch.Tensor) – [1, total_tokens, H, 128] or [total_tokens, H, 128], bfloat16 or float16.

  • g (torch.Tensor) – Channel-wise gate [..., total_tokens, HV, 128]. Log-space unless use_gate_in_kernel, in which case it is the raw pre-activation and cuDNN applies the safe-gate transform from A_log / dt_bias / lower_bound. float32, bfloat16 or float16; cuDNN takes all three and only the gate’s memory format follows the choice, so this is forwarded with no copy. In float16 the kernel’s chunk-cumulative decay inverse bounds how strong the decay may be (roughly alpha >= 0.9 per token per channel before it overflows); bfloat16 carries an fp32-like exponent and has no such bound.

  • beta (torch.Tensor) – [..., total_tokens, HV]. Post-sigmoid in float32 or q.dtype, or q.dtype logits when beta_is_logit.

  • A_log (torch.Tensor, optional) – Safe-gate parameters, required together when use_gate_in_kernel.

  • dt_bias (torch.Tensor, optional) – Safe-gate parameters, required together when use_gate_in_kernel.

  • scale (float, optional) – Query scale; 1 / sqrt(head_dim) when None.

  • output_final_state (bool) – Return the final state alongside the output. This gates only the return value; see initial_state for when a state is written.

  • use_qk_l2norm_in_kernel (bool) – Fuse the q/k L2 normalization into the kernel.

  • use_gate_in_kernel (bool) – Read g as the raw pre-activation and apply the safe-gate transform from A_log / dt_bias / lower_bound in the kernel.

  • beta_is_logit (bool) – Read beta as logits and apply the sigmoid in the kernel.

  • lower_bound (float, optional) – Safe-gate lower bound, forwarded as cuDNN’s gate_lower_bound.

  • cu_seqlens (torch.Tensor) – [num_seqs + 1] int32 or int64. Required.

  • initial_state (torch.Tensor, optional) – State [num_seqs, HV, 128, 128], V-major, float32 or bfloat16. Following the Cake and CuTe DSL prefill backends, initial_state is advanced to the final state whenever one is given and no separate output_state is supplied, independently of output_final_state – which gates only what is returned. output_state must not alias initial_state; see cudnn_chunk_gated_delta_rule().

  • output_state (torch.Tensor, optional) – State [num_seqs, HV, 128, 128], V-major, float32 or bfloat16. Following the Cake and CuTe DSL prefill backends, initial_state is advanced to the final state whenever one is given and no separate output_state is supplied, independently of output_final_state – which gates only what is returned. output_state must not alias initial_state; see cudnn_chunk_gated_delta_rule().

  • output (torch.Tensor, optional) – Pre-allocated output, written in place by the kernel.

  • batch_invariant (bool) – Disable the split-K partition; see cudnn_chunk_gated_delta_rule().

Returns:

(output, final_state), with final_state None when output_final_state=False.

Return type:

Tuple[torch.Tensor, Optional[torch.Tensor]]