flashinfer.kda.recurrent_kda¶
- flashinfer.kda.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, ssm_state_indices: Tensor | None = None, num_spec_tokens: int | None = None, num_accepted_tokens: Tensor | None = None, output: Tensor | None = None, initial_state_source: Tensor | None = None, initial_state_indices: Tensor | None = None, beta_is_logit: bool = False, seq_order: Tensor | None = None, prefill_workspace: RecurrentKDAPrefillWorkspace | None = None, state_checkpoints: Tensor | None = None, checkpoint_cu_starts: Tensor | None = None, checkpoint_every_n_tokens: int = 0, checkpoint_state_indices: Tensor | None = None, *, disable_state_update: bool = False, correction_cache: Tensor | None = None, kg_cache: Tensor | None = None, backend: Literal['auto', 'cute-dsl', 'cake'] = 'auto') tuple[Tensor, Tensor | None] | tuple[Tensor, Tensor | None, Tensor]¶
Recurrent KDA (Kimi Delta Attention) decode and prefill kernel.
This is the public API layer for the CuTe DSL implementation in
flashinfer.kda_kernels.recurrent_kda. It supports single-token decode, fused speculative decode, GQA, optional cu_seqlens packing, and the same gate modes as the backend implementation. On SM120a, eligible ordinary multi-token prefill uses the architecture-specific CuTe DSL backend. On SM100a (B200/GB200) and SM103a (B300/GB300), the FlashKDA-compatible subset can use either the frozen Cake schedules or the source-level CuTe DSL BT=16 kernel. The Cake backend includes a generated two-stage BT=16 prepare/chain portfolio with device- and shape-specific S7/S8/S9 pipeline selection.backend="auto"prefers CuTe DSL for supported plain prefill contracts and keeps Cake as the feature-complete fallback; usebackend="cake"to select and benchmark the generated portfolio explicitly. Compatible equal-head D128 unbounded-softplus T=1 decode calls use their frozen Cake specialization automatically. The Cake path accepts any positive runtime head count, so Kimi-Linear tensor parallelism maps global H32 to per-rank H32/H16/H8/H4 without an adapter. Other decode and speculative-decode calls retain the CuTe DSL backend.- Parameters:
q (torch.Tensor) – Query of shape
[B, T, H, K], or[1, total_tokens, H, K]when usingcu_seqlens. Must be bfloat16.T=1selects decode; eligibleT>1calls may select the frozen prefill backend.k (torch.Tensor) – Key with the same shape as
q. Must be bfloat16.v (torch.Tensor) – Value of shape
[B, T, HV, V], or[1, total_tokens, HV, V]when packed. Must be bfloat16. GQA is applied whenHV != H.g (torch.Tensor) – Per-K-dimension gate of shape
[B, T, HV, K], or[1, total_tokens, HV, K]when packed. Must be bfloat16. Log-space if pre-computed, raw input ifuse_gate_in_kernel=True.beta (torch.Tensor) – Delta-rule learning rate of shape
[B, T, HV], or[1, total_tokens, HV]when packed. Must be bfloat16. Pre-sigmoided unlessbeta_is_logit=True. Eligible frozen prefill accepts non-overlapping token-row-strided storage with a unit head stride, including a view into a fused projection.A_log (Optional[torch.Tensor]) – Log decay parameter of shape
[H]. Must be float32. Required whenuse_gate_in_kernel=True.dt_bias (Optional[torch.Tensor]) – Per-head-K decay bias of shape
[H*K]or[H, K]. Must be float32.scale (Optional[float]) – Scale factor for queries. If
None, defaults to1 / sqrt(K).initial_state (Optional[torch.Tensor]) – Initial state of shape
[N, HV, V, K]. Eligible prefill backends accept bfloat16 or float32; other modes may be stricter. IfNone, zero-initialized. Updated in-place. For batched spec decode withoutcu_seqlens,Nis the packed checkpoint-slot countB * (1 + num_spec_tokens)whenssm_state_indicesis omitted. For eligible frozen prefill withssm_state_indices, this is a state pool[N_pool, H, 128, 128]whose inner slots are contiguous; padding between pool slots is allowed.output_final_state (bool) – Whether to return the final state. Default:
False.use_qk_l2norm_in_kernel (bool) – Whether to apply L2 normalization to Q and K. Default:
True.use_gate_in_kernel (bool) – Whether to compute the gate inside the kernel from
A_logandg. Default:False.lower_bound (Optional[float]) – If set, uses
lower_bound * sigmoid(exp(A_log) * (g + dt_bias))gate formula. IfNone, uses-exp(A_log) * softplus(g + dt_bias). A supplied bound must be negative.cu_seqlens (Optional[torch.Tensor]) – Contiguous CUDA cumulative sequence lengths of shape
[N+1]. May be int32 or int64. Frozen prefill converts int32 offsets to int64 outside graph capture; graph capture requires caller-provided int64 offsets. For frozen prefill, values must start at zero, be non-decreasing, and end at the total token count. This value contract is not normally host-validated. Cake may read these values to prepare and cache host scheduling metadata. CuTe DSL generates packed sequence ordering on the device and can generate decomposed chunk metadata there when usingRecurrentKDAPrefillWrapper. Eligible 148-SM B200 and 152-SM GB200 Cake calls additionally cache persistent worker task bins.ssm_state_indices (Optional[torch.Tensor]) – State cache indices. Shape
[N]int32 for standard decode, or[N, 1+S]int32 for spec decode (num_spec_tokensmust also be set). Eligible frozen packed prefill accepts contiguous CUDA int32[N_seq]indices and updates the selectedinitial_statepool slots directly.num_spec_tokens (Optional[int]) – Number of speculative tokens (S). When set, processes 1+S tokens in a single fused kernel launch. Must be >= 1.
num_accepted_tokens (Optional[torch.Tensor]) – Per-sequence accepted token count from the previous spec decode round. Shape
[N]int32. IfNone, initial state is loaded fromssm_state_indices[n, 0]. Values above1+Sare clamped to the final checkpoint slot.output (Optional[torch.Tensor]) – Pre-allocated output tensor. Shape
[B, T, HV, V]for fixed layout, or the corresponding packed/speculative shape when usingcu_seqlens. IfNone, a new tensor is allocated. Frozen prefill requires storage disjoint from Q, K, V, G, beta, andinitial_state.initial_state_source (Optional[torch.Tensor]) – Optional read-only committed state pool
[N0, HV, V, K]. When provided, token 0 is loaded from this pool instead ofinitial_state.initial_state_indices (Optional[torch.Tensor]) – Source slot per sequence, shape
[N]int32. Required together withinitial_state_source.beta_is_logit (bool) – If
True, apply sigmoid tobetainside the recurrent kernel.disable_state_update (bool) – Frozen / speculative-verify mode: compute outputs for up to 16 tokens per sequence from the committed state and never write any state back (
final_stateisNone). Seeflashinfer.kda_decode.recurrent_kda()for the full mode contract, including the optional slot-indexedcorrection_cache/kg_cacheverify outputs.correction_cache (Optional[torch.Tensor]) – Frozen-verify only: slot-indexed float32 per-token delta-rule corrections
[num_slots, HV, T_max, V].kg_cache (Optional[torch.Tensor]) – Frozen-verify only: slot-indexed bf16 (raw key | raw gate) cache
[num_slots, HV, T_max, 2*K].seq_order (Optional[torch.Tensor]) – Optional packed-prefill sequence order, as a contiguous CUDA int32 permutation of shape
[N]. CuTe DSL packed engine calls generate a stable longest-sequence-first order on the device when this is omitted; CuTe DSL decomp keeps the original order unless the graph wrapper requests device-generated metadata. Cake constructs and caches its own eager host metadata. On Cake, supplying an order disables persistent host task-bin planning but does not force direct M128; the selected non-persistent route may still be BT16 prepare/chain, M64, small-BH, or direct according to the input shape. Fixed-layout prefill and decode calls must leave it asNone.prefill_workspace (Optional[RecurrentKDAPrefillWorkspace]) – Caller-owned workspace for SM100-family and SM120 prefill backends. It is optional for eager execution and required for CUDA graph capture. Warm it eagerly with the exact tensors on the capture stream before capture. Use one workspace per captured
recurrent_kdainvocation. Explicit workspaces and CUDA Graph capture use non-persistent schedules, including eligible BT16, M64, small-BH, and direct routes. Persistent task planning is an eager-only B200/GB200 route because its bins depend on host-visible sequence lengths.state_checkpoints (Optional[torch.Tensor]) – Checkpoint output or pool
[C, H, 128, 128]for prefill. CuTe DSL accepts BF16 or FP32 and requires it to matchinitial_statewhen present; Cake accepts BF16. Withoutcheckpoint_state_indices, row zero for each sequence is its initial state and later rows are states before token blocks beginning atN, 2N, .... CuTe DSL allocates this packed output during eager execution when omitted; CUDA graph capture requires a caller-owned tensor.checkpoint_cu_starts (Optional[torch.Tensor]) – Contiguous CUDA int64 cumulative checkpoint counts
[N_seq+1]. The first value must be zero. Withoutcheckpoint_state_indices, each difference isceil(seq_len / N). With indices, each difference isfloor(seq_len / N)and counts completed periodic boundaries, including an aligned final boundary and excluding the initial state.checkpoint_state_indices (Optional[torch.Tensor]) – CuTe DSL-only contiguous CUDA int32 destination rows. Entry
iselects the row of thestate_checkpointspool written for packed completed-boundary entryi. The kernel writes the pool directly; no temporary checkpoint tensor or scatter is used.checkpoint_every_n_tokens (int) – Checkpoint interval. Zero disables checkpoints; a positive value must be divisible by 32, except that the SM100-family exact-N16 frozen route also accepts multiples of 16. SGLang normally uses 64 or a larger cache-page-aligned multiple.
backend (Literal["auto", "cute-dsl", "cake"]) – Implementation backend.
"auto"selects the architecture- appropriate CuTe DSL kernel for supported ordinary multi-token prefill, including the SM120 backend and SM100-family state checkpoints, and otherwise falls back to an exported frozen Cake specialization."cake"and"cute-dsl"select those backends strictly. The Cake prefill path chooses among direct, persistent, small-BH, and two-stage BT16 schedules from the input shape and physical device. The SM100-family kernel additionally needsnvidia-cutlass-dsl>=4.7; below that"auto"uses Cake there and"cute-dsl"raisesImportError.
- Returns:
Tuple of
(output, final_state)wherefinal_stateisNonewhenoutput_final_state=False. When checkpointing is enabled, a triple(output, final_state, state_checkpoints)is returned. Seeflashinfer.kda_kernels.recurrent_kda.run_recurrent_kda()for the backend implementation.