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_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, *, write_state=True)

source

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.
  • write_state (bool) – Whether to write the updated window back to conv_state. False only reads it for the look-back.

Returns:

The conv output, [total_seq_len, conv_dim].

Return type:

TensorValue