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_conv1d_fwd
gated_delta_conv1d_fwd()
max.nn.state_space.gated_delta_conv1d_fwd(qkv_input_ragged, conv_weight, conv_state, slot_idx, input_row_offsets)
Applies the causal conv1d pass, mutating a slot-indexed conv-state pool in place.
conv_state is a mutable pool of shape [max_slots, conv_dim, kernel_size - 1] and the kernel reads/writes slot
slot_idx[batch_item] directly. There is no conv_state_out
graph output: the pool is mutated in place.
-
Parameters:
-
- qkv_input_ragged (TensorValue) – The
[total_seq_len, conv_dim]projected QKV input. - conv_weight (TensorValue) – The
[conv_dim, kernel_size]depthwise conv weights. - conv_state (BufferValue) – The
[max_slots, conv_dim, kernel_size - 1]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_input_ragged (TensorValue) – The
-
Returns:
-
The conv output,
[total_seq_len, conv_dim]. -
Return type: