IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo function

causal_conv1d_update_cpu

def causal_conv1d_update_cpu[x_dtype: DType, conv_state_dtype: DType, weight_dtype: DType, output_dtype: DType, bias_dtype: DType](batch: Int, dim: Int, seqlen: Int, width: Int, state_len: Int, x: TileTensor[x_dtype, Storage=x.Storage, address_space=x.address_space, linear_idx_type=x.linear_idx_type], conv_state: TileTensor[conv_state_dtype, Storage=conv_state.Storage, address_space=conv_state.address_space, linear_idx_type=conv_state.linear_idx_type], weight: TileTensor[weight_dtype, Storage=weight.Storage, address_space=weight.address_space, linear_idx_type=weight.linear_idx_type], output: TileTensor[output_dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], bias: TileTensor[bias_dtype, Storage=bias.Storage, address_space=bias.address_space, linear_idx_type=bias.linear_idx_type], x_batch_stride: UInt32, x_c_stride: UInt32, x_l_stride: UInt32, conv_state_batch_stride: UInt32, conv_state_c_stride: UInt32, conv_state_l_stride: UInt32, weight_c_stride: UInt32, weight_width_stride: UInt32, out_batch_stride: UInt32, out_c_stride: UInt32, out_l_stride: UInt32, silu_activation: Bool)

CPU implementation of causal conv1d update for incremental inference.

This kernel:

  1. Concatenates conv_state with x to form a sliding window
  2. Computes convolution output for the new positions
  3. Updates conv_state with the new values from x

Simple mode (no circular buffer):

  • conv_state holds the last (state_len) values
  • New x values are appended, old values are shifted out

Parameters:

  • ​x_dtype (DType): Element type of the input tensor x.
  • ​conv_state_dtype (DType): Element type of the convolution state tensor conv_state.
  • ​weight_dtype (DType): Element type of the weight tensor weight.
  • ​output_dtype (DType): Element type of the output tensor output.
  • ​bias_dtype (DType): Element type of the bias tensor bias.

Args:

Was this page helpful?