IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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')

source

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, before dt_bias is 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 per state_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_BOUND before 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.

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:

TensorValue