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).
Mojo function
gated_delta_recurrence_fwd_shape
def gated_delta_recurrence_fwd_shape(qkv_conv_output: T, decay_per_token: T, beta_per_token: T, recurrent_state: T, slot_idx: T, input_row_offsets: T) -> IndexList[Int(2)]
Computes the output shape for the gated_delta_recurrence_fwd graph op.
Args:
- qkv_conv_output (
T): Ragged conv output of shape[total_seq_len, conv_dim]. - decay_per_token (
T): Per-token decay factors of shape[total_seq_len, num_value_heads]. - beta_per_token (
T): Per-token beta gates of shape[total_seq_len, num_value_heads]. - recurrent_state (
T): Mutable recurrent-state pool of shape[max_slots, num_value_heads, key_head_dim, value_head_dim]. - slot_idx (
T): Per-batch slot indices into the recurrent-state pool, shape[batch_size]. - input_row_offsets (
T): Cumulative row offsets per batch, shape[batch_size + 1].
Returns: