flashinfer.kda_decode.packed_kda_decode¶
- flashinfer.kda_decode.packed_kda_decode(mixed_qkv: Tensor, raw_gate: Tensor, raw_beta: Tensor, A_log: Tensor, dt_bias: Tensor, state: Tensor, state_indices: Tensor, output: Tensor | None = None) Tensor¶
Run serving-native packed Kimi K3 recurrent decode.
This operator consumes the post-convolution packed QKV row and raw gate and beta logits directly. It fuses Q/K extraction and L2 normalization, the Kimi K3 lower-bound gate transform, beta sigmoid, and one recurrent state update into a single exported Cake kernel. It is specialized for
T=1,H=12, andK=V=128on exact SM100a and SM103a devices.The fixed numerical contract uses
scale=1/sqrt(128), L2 epsilon1e-6, andlower_bound=-5.stateis updated in place on the caller’s current PyTorch CUDA stream. Batches below 32 use the eight-row value tile; batches of 32 or more use the sixteen-row value tile.- Parameters:
mixed_qkv – Post-convolution packed QKV with shape
[B, 3 * 12 * 128]and dtype bfloat16. The last dimension must be contiguous; positive padding between batch rows is allowed.raw_gate – Raw per-channel recurrence gate with shape
[B, 12 * 128]and dtype bfloat16. The last dimension must be contiguous.raw_beta – Raw delta-rule learning-rate logits with shape
[B, 12]and dtype bfloat16. The last dimension must be contiguous.A_log – Contiguous float32 log-decay parameter with shape
[12].dt_bias – Contiguous float32 per-channel decay bias with shape
[12 * 128].state – Caller-owned bfloat16 recurrent-state pool with shape
[N, 12, 128, 128]. Its inner three dimensions must be compact; the outer slot stride may contain arbitrary positive padding.state_indices – Contiguous CUDA int32 cache slot for each row, with shape
[B]. Active indices must be unique and in bounds.-1marks an inactive CUDA-graph padding row, which produces zero output and does not access or updatestate. These value constraints are not host-validated, avoiding a device synchronization.output – Optional caller-owned contiguous bfloat16 output with shape
[B, 1, 12, 128]. Supplying it avoids allocation and is required for an allocation-free CUDA-graph replay path.
- Returns:
The bfloat16 output with shape
[B, 1, 12, 128]by default. Whenoutputis supplied, the returned tensor is that exact allocation.