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)
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 ofgated_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.
- qkv_conv_output (TensorValue) – The
-
Returns:
-
The recurrence output,
[total_seq_len, value_dim]. -
Return type: