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 struct
GatedDeltaStateFold
struct GatedDeltaStateFold
Applies a speculative 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_fwd, for every layer in
one launch. A row with num_accepted[b] == 0 is left unchanged.
The verify width K is the length of verify_width, whose contents are
never read. At K == 0 the verify wrote no records and landed the state
itself, so nothing is launched.
Tensor Shapes: - recurrent_state : [max_slots, num_value_heads, KD, VD] (MUT) - ring : [ring_rows, num_key_heads, RING_LEN, record_stride] - row_ids : [num_layers, batch_size] uint32 - ring_row_ids : [num_layers, batch_size] uint32 - num_accepted : [batch_size] uint32 - verify_width : [K] int64
Implemented traitsโ
Methodsโ
executeโ
static def execute[state_dtype: DType, ring_dtype: DType, target: StringSpan[ImmStaticOrigin]](recurrent_state: ManagedTensorSlice[IOSpec[_, _].MutableInput, static_spec=recurrent_state.static_spec], ring: ManagedTensorSlice[IOSpec[_, _].MutableInput, static_spec=ring.static_spec], row_ids: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=row_ids.static_spec], ring_row_ids: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=ring_row_ids.static_spec], num_accepted: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=num_accepted.static_spec], verify_width: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=verify_width.static_spec], ctx: DeviceContext)