flashinfer.kda_decode.packed_fused_kda_decode

flashinfer.kda_decode.packed_fused_kda_decode(x: Tensor, weight: Tensor, conv_state: Tensor, raw_gate: Tensor, raw_beta: Tensor, A_log: Tensor, dt_bias: Tensor, state_indices: Tensor, state: Tensor, output_gate: Tensor, norm_weight: Tensor, lower_bound: float | None = -5.0, norm_eps: float = 1e-05, output: Tensor | None = None, *, query_start_loc: Tensor, num_accepted_tokens: Tensor, t1_state_indices: Tensor | None = None) → Tensor

Run packed T>=1 fused KDA with per-token cache checkpoints.

Tokens use a packed ragged layout described by query_start_loc. state_indices has shape [N, T], where T is the maximum verification length. num_accepted_tokens[n] - 1 selects the source checkpoint and convolution-history offset for sequence n; token t writes its recurrent checkpoint to state_indices[n, t]. The convolution cache is one rolling window of length T + 2 in state_indices[n, 0].

This facade is CuTe-only: T=1 uses the CuTe DSL implementation underlying fused_kda_decode (without Cake or backend selection), including bfloat16-state and softplus-gate support. The optional t1_state_indices optimization is strictly T=1-only; passing it for T>1 fails closed. T>1 uses the SM10x CuTe DSL packed backend, which accepts float32 or bfloat16 recurrent state and requires a finite negative lower_bound. Both paths require head dimension 128 and convolution width four; T=1 accepts 8, 12, 24, 32, 48, or 96 heads and T>1 accepts 12, 24, 32, 48, or 96 heads. Other tensors follow fused_kda_decode(), with the packed row count supplied by x.

query_start_loc must be contiguous int32 of shape [N+1], start at zero, and contain nondecreasing offsets within x with active lengths no larger than T. num_accepted_tokens must be contiguous int32 of shape [N] with values in [1, T]. For T=1, t1_state_indices may optionally provide the already-resolved contiguous per-row cache slots with shape [num_rows]. When supplied, it is used by the direct fused T=1 specialization and query_start_loc is retained only as structural metadata; its values are not read by the host or kernel. This is intended for CUDA-Graph-stable caller-owned buffers updated in-place. Active recurrent destinations must be positive cache slots and must not alias destinations from another active sequence; non-positive or zero-length rows are null rows and do not mutate either cache. In T=1, packed rows beyond query_start_loc[-1] are trailing unused capacity: they receive zero output and do not mutate either cache. These value constraints are caller-owned so CUDA Graph replay never performs a device-to-host validation sync.

Parameters:
  • x – Packed QKV projection with shape [num_rows, 3 * H * 128] and dtype bfloat16.

  • weight – Depthwise convolution weights with shape [3, 4, H * 128] and dtype float32.

  • conv_state – Paged bfloat16 convolution cache. T=1 uses history length three; T>1 uses the extended rolling history length T + 2.

  • raw_gate – Raw recurrence gate with shape [1, num_rows, H, 128] and dtype bfloat16.

  • raw_beta – Raw delta-rule learning-rate logits with shape [1, num_rows, H] and dtype bfloat16.

  • A_log – Log decay parameter with H elements and dtype float32.

  • dt_bias – Per-channel decay bias with H * 128 elements and dtype float32.

  • state_indices – Contiguous int32 cache slots with shape [N, T]. For T>1, column t is the recurrent checkpoint destination for token t. Column zero also selects the rolling convolution cache.

  • query_start_loc – Contiguous int32 packed-row offsets with shape [N + 1].

  • num_accepted_tokens – Contiguous int32 accepted-token counts with shape [N].

  • t1_state_indices – Optional contiguous int32 per-packed-row cache slots with shape [num_rows]. T=1 only; null and trailing rows must be zero or non-positive. The tensor is never allocated or populated by this function.

  • state – Paged recurrent state with shape [num_slots, H, 128, 128]. float32 or bfloat16. T>1 keeps the recurrence in float32 and rounds only when writing each bfloat16 checkpoint.

  • output_gate – Gated RMSNorm logits with shape [num_rows, H, 128] or [1, num_rows, H, 128] and dtype bfloat16.

  • norm_weight – RMSNorm weight with 128 elements and dtype float32.

  • lower_bound – Negative recurrence-gate lower bound. T=1 also accepts None for the softplus gate; T>1 requires a finite negative value.

  • norm_eps – Non-negative RMSNorm epsilon.

  • output – Optional preallocated contiguous bfloat16 output with shape [1, num_rows, H, 128].

Returns:

The packed bfloat16 output with shape [1, num_rows, H, 128].