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):DTypeof the per-token input and output tensors. - state_dtype (
DType):DTypeof therecurrent_statepool. - ring_dtype (
DType):DTypeof the ring pool. The fold is bit-exact against a forward over the accepted prefix only forfloat32. - 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. - 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. - ring_LT (
TensorLayout):TensorLayoutforring. - ring_slot_idx_LT (
TensorLayout):TensorLayoutforring_slot_idx. - Engine (
TensorEngine): Engine shared by all tile operands.
Args:
- batch_size (
Int32): Number of sequences in the ragged batch. - num_value_heads (
Int32): Number of value heads (nv). - num_key_heads (
Int32): Number of key heads (nk). - key_dim (
Int32):num_key_heads * key_head_dim. - recurrence_output (
TileTensor[work_dtype, recurrence_output_LT, MutUntrackedOrigin, Engine=Engine]):[total_seq_len, value_dim]readout, written for every token of the window. - recurrent_state (
TileTensor[state_dtype, recurrent_state_LT, MutUntrackedOrigin, Engine=Engine]):[max_slots, nv, KEY_HEAD_DIM, VALUE_HEAD_DIM]dense live pool, written only for a row longer thanRING_LEN. - slot_idx (
TileTensor[.uint32, slot_idx_LT, MutUntrackedOrigin, Engine=Engine]):[batch_size]live pool row per batch item. - qkv_conv_output (
TileTensor[work_dtype, qkv_conv_output_LT, MutUntrackedOrigin, Engine=Engine]):[total_seq_len, conv_dim]conv output. - decay_per_token (
TileTensor[work_dtype, decay_per_token_LT, MutUntrackedOrigin, Engine=Engine]):[total_seq_len, nv]per-token decay. - beta_per_token (
TileTensor[work_dtype, beta_per_token_LT, MutUntrackedOrigin, Engine=Engine]):[total_seq_len, nv]per-token beta gate. - input_row_offsets (
TileTensor[.uint32, input_row_offsets_LT, MutUntrackedOrigin, Engine=Engine]):[batch_size + 1]ragged offsets. - ring (
TileTensor[ring_dtype, ring_LT, MutUntrackedOrigin, Engine=Engine]):[ring_rows, nk, RING_LEN, ring_record_stride]dense ring pool. - ring_slot_idx (
TileTensor[.uint32, ring_slot_idx_LT, MutUntrackedOrigin, Engine=Engine]):[batch_size]ring row per batch item. - ring_record_stride (
Int32): Elements between consecutive records, at leastgated_delta_ring_record_elements(nv // nk). - qkv_conv_output_seqlen_stride (
UInt32): Stride between sequence positions inqkv_conv_output. - qkv_conv_output_channel_stride (
UInt32): Stride between channels inqkv_conv_output. - per_token_seqlen_stride (
UInt32): Stride between sequence positions indecay_per_tokenandbeta_per_token. - per_token_head_stride (
UInt32): Stride between heads indecay_per_tokenandbeta_per_token. - recurrence_output_seqlen_stride (
UInt32): Stride between sequence positions inrecurrence_output. - recurrence_output_valuedim_stride (
UInt32): Stride between value-dim elements inrecurrence_output.