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).

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, Engine: TensorEngine](batch_size: Int32, num_value_heads: Int32, num_key_heads: Int32, key_dim: Int32, recurrence_output: TileTensor[work_dtype, recurrence_output_LT, MutUntrackedOrigin, Engine=Engine], recurrent_state: TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin, Engine=Engine], slot_idx: TileTensor[.uint32, slot_idx_LT, MutUntrackedOrigin, Engine=Engine], qkv_conv_output: TileTensor[work_dtype, qkv_conv_output_LT, MutUntrackedOrigin, Engine=Engine], decay_per_token: TileTensor[work_dtype, decay_per_token_LT, MutUntrackedOrigin, Engine=Engine], beta_per_token: TileTensor[work_dtype, beta_per_token_LT, MutUntrackedOrigin, Engine=Engine], input_row_offsets: TileTensor[.uint32, input_row_offsets_LT, MutUntrackedOrigin, Engine=Engine], qkv_conv_output_seqlen_stride: UInt32, qkv_conv_output_channel_stride: UInt32, per_token_seqlen_stride: UInt32, per_token_head_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): DType for the per-token input and output tensors (qkv_conv_output, decay_per_token, beta_per_token, recurrence_output), float32.
  • ​state_dtype (DType): DType for the recurrent_state pool (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 equal KEY_HEAD_DIM.
  • ​recurrence_output_LT (TensorLayout): TensorLayout for recurrence_output.
  • ​qkv_conv_output_LT (TensorLayout): TensorLayout for qkv_conv_output.
  • ​decay_per_token_LT (TensorLayout): TensorLayout for decay_per_token.
  • ​beta_per_token_LT (TensorLayout): TensorLayout for beta_per_token.
  • ​recurrent_state_LT (TensorLayout): TensorLayout for recurrent_state.
  • ​slot_idx_LT (TensorLayout): TensorLayout for slot_idx.
  • ​input_row_offsets_LT (TensorLayout): TensorLayout for input_row_offsets.
  • ​Engine (TensorEngine): Engine shared by all tile operands.

Args:

Was this page helpful?