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, noinitial_state_source, no state checkpoints.Requires cudnn-frontend 1.29+ with the
cutedslextra. 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 unlessuse_gate_in_kernel, in which case it is the raw pre-activation and cuDNN applies the safe-gate transform fromA_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 (roughlyalpha >= 0.9per 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 orq.dtype, orq.dtypelogits whenbeta_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)whenNone.output_final_state (bool) – Return the final state alongside the output. This gates only the return value; see
initial_statefor 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
gas the raw pre-activation and apply the safe-gate transform fromA_log/dt_bias/lower_boundin the kernel.beta_is_logit (bool) – Read
betaas 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_stateis advanced to the final state whenever one is given and no separateoutput_stateis supplied, independently ofoutput_final_state– which gates only what is returned.output_statemust not aliasinitial_state; seecudnn_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_stateis advanced to the final state whenever one is given and no separateoutput_stateis supplied, independently ofoutput_final_state– which gates only what is returned.output_statemust not aliasinitial_state; seecudnn_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), withfinal_stateNonewhenoutput_final_state=False.- Return type:
Tuple[torch.Tensor, Optional[torch.Tensor]]