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
gated_delta_recurrence_fwd_gpu
def gated_delta_recurrence_fwd_gpu[work_dtype: DType, state_dtype: DType, KEY_HEAD_DIM: Int, VALUE_HEAD_DIM: Int, recurrence_output_LT: TensorLayout, qkv_conv_output_LT: TensorLayout, decay_per_token_LT: TensorLayout, beta_per_token_LT: TensorLayout, recurrent_state_LT: TensorLayout, slot_idx_LT: TensorLayout, input_row_offsets_LT: TensorLayout](batch_size: Int, num_value_heads: Int, num_key_heads: Int, key_dim: Int, recurrence_output: TileTensor[work_dtype, recurrence_output_LT, MutUntrackedOrigin], recurrent_state: TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin], slot_idx: TileTensor[DType.uint32, slot_idx_LT, MutUntrackedOrigin], qkv_conv_output: TileTensor[work_dtype, qkv_conv_output_LT, MutUntrackedOrigin], decay_per_token: TileTensor[work_dtype, decay_per_token_LT, MutUntrackedOrigin], beta_per_token: TileTensor[work_dtype, beta_per_token_LT, MutUntrackedOrigin], input_row_offsets: TileTensor[DType.uint32, input_row_offsets_LT, MutUntrackedOrigin], qkv_conv_output_seqlen_stride: UInt32, qkv_conv_output_channel_stride: UInt32, per_token_seqlen_stride: UInt32, per_token_head_stride: UInt32, recurrent_state_slot_stride: UInt32, recurrent_state_value_head_stride: UInt32, recurrent_state_key_dim_stride: UInt32, recurrent_state_value_dim_stride: UInt32, recurrence_output_seqlen_stride: UInt32, recurrence_output_valuedim_stride: UInt32)
GPU kernel: slot-indexed gated delta rule recurrence, one CTA per head.
One CTA owns one (batch_item, value_head); thread tid == vd_element owns
the KD-element state column recurrent_state[slot, value_head, :, tid] in
registers for the whole sequence. The per-token raw Q/K for this value
head's key head are staged once per block in shared memory (one element per
thread, coalesced) so the KD reductions read them from shared memory rather
than every vd-thread re-reading the same KD elements from global memory;
L2 normalisation and the 1/sqrt(KD) query scale are folded in as scalars
factored out of the reductions.
Parameters:
- βwork_dtype (
DType):DTypefor the per-token input and output tensors (qkv_conv_output,decay_per_token,beta_per_token,recurrence_output),float32. - βstate_dtype (
DType):DTypefor therecurrent_statepool (bfloat16). - βKEY_HEAD_DIM (
Int): Compile-time key head dimension (e.g. 128 for Qwen3.5). - βVALUE_HEAD_DIM (
Int): Compile-time value head dimension; must equalKEY_HEAD_DIM. - βrecurrence_output_LT (
TensorLayout):TensorLayoutforrecurrence_output. - βqkv_conv_output_LT (
TensorLayout):TensorLayoutforqkv_conv_output. - βdecay_per_token_LT (
TensorLayout):TensorLayoutfordecay_per_token. - βbeta_per_token_LT (
TensorLayout):TensorLayoutforbeta_per_token. - βrecurrent_state_LT (
TensorLayout):TensorLayoutforrecurrent_state. - βslot_idx_LT (
TensorLayout):TensorLayoutforslot_idx. - βinput_row_offsets_LT (
TensorLayout):TensorLayoutforinput_row_offsets.
Args:
- βbatch_size (
Int): Number of sequences in the ragged batch. - βnum_value_heads (
Int): Number of value heads (nv). - βnum_key_heads (
Int): Number of key heads (nk); the GQA expansion ratio isnum_value_heads / num_key_heads. - βkey_dim (
Int): Total key dimension, equal tonum_key_heads * KEY_HEAD_DIM. - βrecurrence_output (
TileTensor[work_dtype, recurrence_output_LT, MutUntrackedOrigin]): Flat output of shape[total_seq_len, value_dim]holding the recurrence result for every token. - βrecurrent_state (
TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin]): Mutable state pool of shape[max_slots, num_value_heads, KEY_HEAD_DIM, VALUE_HEAD_DIM]; the kernel reads and writes slotslot_idx[batch_item]in place. - βslot_idx (
TileTensor[DType.uint32, slot_idx_LT, MutUntrackedOrigin]): Pool slot index for each batch item, shape[batch_size],uint32. - βqkv_conv_output (
TileTensor[work_dtype, qkv_conv_output_LT, MutUntrackedOrigin]): Conv1d output from Pass 1, shape[total_seq_len, conv_dim], with Q in channels[0, key_dim), K in[key_dim, 2*key_dim), V in[2*key_dim, 2*key_dim + value_dim). - βdecay_per_token (
TileTensor[work_dtype, decay_per_token_LT, MutUntrackedOrigin]): Per-token per-head scalar decay factor, shape[total_seq_len, num_value_heads]. - βbeta_per_token (
TileTensor[work_dtype, beta_per_token_LT, MutUntrackedOrigin]): Per-token per-head beta gate, shape[total_seq_len, num_value_heads]. - βinput_row_offsets (
TileTensor[DType.uint32, input_row_offsets_LT, MutUntrackedOrigin]): Ragged offsets of shape[batch_size + 1]; sequencebspans flat indices[input_row_offsets[b], input_row_offsets[b+1]). - βqkv_conv_output_seqlen_stride (
UInt32): Stride between consecutive sequence positions inqkv_conv_output. - βqkv_conv_output_channel_stride (
UInt32): Stride between consecutive channels inqkv_conv_output. - βper_token_seqlen_stride (
UInt32): Stride between consecutive sequence positions indecay_per_tokenandbeta_per_token. - βper_token_head_stride (
UInt32): Stride between consecutive heads indecay_per_tokenandbeta_per_token. - βrecurrent_state_slot_stride (
UInt32): Stride between consecutive slots inrecurrent_state. - βrecurrent_state_value_head_stride (
UInt32): Stride between consecutive value heads inrecurrent_state. - βrecurrent_state_key_dim_stride (
UInt32): Stride between consecutive key-dim elements inrecurrent_state. - βrecurrent_state_value_dim_stride (
UInt32): Stride between consecutive value-dim elements inrecurrent_state. - βrecurrence_output_seqlen_stride (
UInt32): Stride between consecutive sequence positions inrecurrence_output. - βrecurrence_output_valuedim_stride (
UInt32): Stride between consecutive value-dim elements inrecurrence_output.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!