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_verify_ring_gpu

def gated_delta_recurrence_verify_ring_gpu[work_dtype: DType, state_dtype: DType, ring_dtype: DType, KEY_HEAD_DIM: Int, VALUE_HEAD_DIM: Int, RING_LEN: 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, ring_LT: TensorLayout, ring_slot_idx_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], ring: TileTensor[ring_dtype, ring_LT, MutUntrackedOrigin, Engine=Engine], ring_slot_idx: TileTensor[.uint32, ring_slot_idx_LT, MutUntrackedOrigin, Engine=Engine], ring_record_stride: Int32, 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: runs the gated delta recurrence over a verify window and records a ring instead of writing the state back.

Produces the same recurrence_output as gated_delta_recurrence_fwd_gpu and leaves recurrent_state unchanged. Each token writes the record gated_delta_state_fold_gpu needs to advance the state over it.

A row longer than RING_LEN records nothing and writes its state back as gated_delta_recurrence_fwd_gpu does, committing the whole row. The fold skips a row whose count exceeds RING_LEN, so a caller folds such a row with its full length.

Parameters:

  • ​work_dtype (DType): DType of the per-token input and output tensors.
  • ​state_dtype (DType): DType of the recurrent_state pool.
  • ​ring_dtype (DType): DType of the ring pool. The fold is bit-exact against a forward over the accepted prefix only for float32.
  • ​KEY_HEAD_DIM (Int): Compile-time key head dimension.
  • ​VALUE_HEAD_DIM (Int): Compile-time value head dimension, equal to KEY_HEAD_DIM.
  • ​RING_LEN (Int): Compile-time record capacity of one ring row.
  • ​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.
  • ​ring_LT (TensorLayout): TensorLayout for ring.
  • ​ring_slot_idx_LT (TensorLayout): TensorLayout for ring_slot_idx.
  • ​Engine (TensorEngine): Engine shared by all tile operands.

Args:

Was this page helpful?