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_state_fold_gpu

def gated_delta_state_fold_gpu[state_dtype: DType, ring_dtype: DType, KEY_HEAD_DIM: Int, VALUE_HEAD_DIM: Int, RING_LEN: Int, recurrent_state_LT: TensorLayout, row_ids_LT: TensorLayout, ring_LT: TensorLayout, ring_row_ids_LT: TensorLayout, num_accepted_LT: TensorLayout, Engine: TensorEngine, KEY_DIM_TILE: Int = Int(8)](batch_size: Int32, num_layers: Int32, num_value_heads: Int32, num_key_heads: Int32, recurrent_state: TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin, Engine=Engine], row_ids: TileTensor[.uint32, row_ids_LT, MutUntrackedOrigin, Engine=Engine], ring: TileTensor[ring_dtype, ring_LT, MutUntrackedOrigin, Engine=Engine], ring_row_ids: TileTensor[.uint32, ring_row_ids_LT, MutUntrackedOrigin, Engine=Engine], ring_record_stride: Int32, num_accepted: TileTensor[.uint32, num_accepted_LT, MutUntrackedOrigin, Engine=Engine])

GPU kernel: applies a verify's accepted records to the state pool.

Advances recurrent_state over the first num_accepted[b] records written by gated_delta_recurrence_verify_ring_gpu. The result is bit-exact against running gated_delta_recurrence_fwd_gpu over the accepted tokens. One launch covers every layer, with the layer as a grid coordinate.

A row is not read or written when num_accepted[b] is zero, or when it exceeds RING_LEN, which names a row the verify was too long to record and committed itself. Rows that alias, as padding requests do on the null block, race, and nothing reads such a row.

Parameters:

  • ​state_dtype (DType): DType of the recurrent_state pool.
  • ​ring_dtype (DType): DType of the ring pool.
  • ​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.
  • ​recurrent_state_LT (TensorLayout): TensorLayout for recurrent_state.
  • ​row_ids_LT (TensorLayout): TensorLayout for row_ids.
  • ​ring_LT (TensorLayout): TensorLayout for ring.
  • ​ring_row_ids_LT (TensorLayout): TensorLayout for ring_row_ids.
  • ​num_accepted_LT (TensorLayout): TensorLayout for num_accepted.
  • ​Engine (TensorEngine): Engine shared by all tile operands.
  • ​KEY_DIM_TILE (Int): Key-dim elements a thread holds in registers at once. Must divide KEY_HEAD_DIM.

Args:

Was this page helpful?