For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Python function
kda_decode
kda_decode()
max.nn.state_space.kda_decode(q, k, v, raw_gate, beta_logits, a_log, dt_bias, cu_seqlens, state_pool, state_indices, *, output_dtype, gate_mode='original', beta_mode='logits', state_layout='K_FIRST')
Runs the KDA recurrence, mutating state_pool in place.
q and k are L2-normalized and q scaled by
1 / sqrt(key_head_dim) inside the kernel, so pass them raw. The gate
and beta activations are folded in too, selected by gate_mode and
beta_mode.
-
Parameters:
-
- q (TensorValue) –
[total_tokens, num_key_heads, key_head_dim]. - k (TensorValue) –
[total_tokens, num_key_heads, key_head_dim]. - v (TensorValue) –
[total_tokens, num_value_heads, value_head_dim]. - raw_gate (TensorValue) –
[total_tokens, num_value_heads, key_head_dim]forget-gate pre-activation, beforedt_biasis added. - beta_logits (TensorValue) –
[total_tokens, num_value_heads]. - a_log (TensorValue) –
[num_value_heads]. - dt_bias (TensorValue) –
[num_value_heads, key_head_dim]. - cu_seqlens (TensorValue) –
[batch_size + 1]int32 exclusive prefix offsets. - state_pool (BufferValue) –
[max_slots, num_value_heads, key_head_dim, value_head_dim]mutable pool, laid out perstate_layout. - state_indices (TensorValue) –
[batch_size]int32 pool slot per sequence. - output_dtype (DType) – Dtype of the returned tensor.
- gate_mode (Literal['original', 'safe']) – Forget-gate form; see
KDA_GATE_LOWER_BOUNDbefore selecting"safe". - beta_mode (Literal['logits', 'probability']) – Whether the kernel applies the sigmoid to
beta_logits. - state_layout (Literal['K_FIRST', 'V_FIRST']) – Pool axis order; must match how the pool was allocated.
- q (TensorValue) –
-
Returns:
-
[total_tokens, num_value_heads, value_head_dim]. -
Raises:
-
ValueError – If the dtype groupings bind no kernel, or the ranks and extents disagree. The kernel’s own shape guards are
debug_assert, so they are absent from a production build. -
Return type: