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):DTypeof therecurrent_statepool. - ring_dtype (
DType):DTypeof the ring pool. - KEY_HEAD_DIM (
Int): Compile-time key head dimension. - VALUE_HEAD_DIM (
Int): Compile-time value head dimension, equal toKEY_HEAD_DIM. - RING_LEN (
Int): Compile-time record capacity of one ring row. - recurrent_state_LT (
TensorLayout):TensorLayoutforrecurrent_state. - row_ids_LT (
TensorLayout):TensorLayoutforrow_ids. - ring_LT (
TensorLayout):TensorLayoutforring. - ring_row_ids_LT (
TensorLayout):TensorLayoutforring_row_ids. - num_accepted_LT (
TensorLayout):TensorLayoutfornum_accepted. - Engine (
TensorEngine): Engine shared by all tile operands. - KEY_DIM_TILE (
Int): Key-dim elements a thread holds in registers at once. Must divideKEY_HEAD_DIM.
Args:
- batch_size (
Int32): Number of sequences the verify covered. - num_layers (
Int32): Number of layers, each with its own pool row. - num_value_heads (
Int32): Number of value heads (nv). - num_key_heads (
Int32): Number of key heads (nk). - recurrent_state (
TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin, Engine=Engine]):[rows, nv, KEY_HEAD_DIM, VALUE_HEAD_DIM]dense live pool, folded in place. - row_ids (
TileTensor[.uint32, row_ids_LT, MutUntrackedOrigin, Engine=Engine]):[num_layers, batch_size]live pool row per (layer, request). - ring (
TileTensor[ring_dtype, ring_LT, MutUntrackedOrigin, Engine=Engine]):[ring_rows, nk, RING_LEN, ring_record_stride]dense ring pool. - ring_row_ids (
TileTensor[.uint32, ring_row_ids_LT, MutUntrackedOrigin, Engine=Engine]):[num_layers, batch_size]ring row per (layer, request). - ring_record_stride (
Int32): Elements between consecutive records. - num_accepted (
TileTensor[.uint32, num_accepted_LT, MutUntrackedOrigin, Engine=Engine]):[batch_size]records to fold.