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

gated_delta_recurrence_fwd

gated_delta_recurrence_fwd()​

max.nn.state_space.gated_delta_recurrence_fwd(qkv_conv_output, decay_per_token, beta_per_token, recurrent_state, slot_idx, input_row_offsets)

source

Applies the gated delta recurrence pass, mutating a slot-indexed state pool in place.

recurrent_state is a mutable pool of shape [max_slots, nv, KD, VD] and the kernel reads/writes slot slot_idx[batch_item] directly. There is no recurrent_state_out graph output: the pool is mutated in place.

Parameters:

  • qkv_conv_output (TensorValue) – The [total_seq_len, conv_dim] output of gated_delta_conv1d_fwd().
  • decay_per_token (TensorValue) – The [total_seq_len, num_value_heads] decays.
  • beta_per_token (TensorValue) – The [total_seq_len, num_value_heads] beta gates.
  • recurrent_state (BufferValue) – The [max_slots, nv, KD, VD] mutable pool.
  • slot_idx (TensorValue) – The [batch_size] uint32 slot indices into the pool.
  • input_row_offsets (TensorValue) – The [batch_size + 1] uint32 ragged offsets.

Returns:

The recurrence output, [total_seq_len, value_dim].

Return type:

TensorValue