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)

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.

Returns:

The conv output, [total_seq_len, conv_dim].

Return type:

TensorValue