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_indiceshas shape[N, T], where T is the maximum verification length.num_accepted_tokens[n] - 1selects the source checkpoint and convolution-history offset for sequencen; tokentwrites its recurrent checkpoint tostate_indices[n, t]. The convolution cache is one rolling window of lengthT + 2instate_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 optionalt1_state_indicesoptimization 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 negativelower_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 followfused_kda_decode(), with the packed row count supplied byx.query_start_locmust be contiguous int32 of shape[N+1], start at zero, and contain nondecreasing offsets withinxwith active lengths no larger than T.num_accepted_tokensmust be contiguous int32 of shape[N]with values in[1, T]. For T=1,t1_state_indicesmay 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 andquery_start_locis 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 beyondquery_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
Helements and dtype float32.dt_bias – Per-channel decay bias with
H * 128elements and dtype float32.state_indices – Contiguous int32 cache slots with shape
[N, T]. For T>1, columntis the recurrent checkpoint destination for tokent. 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
Nonefor 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].