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
kpool_seed_tail_kernel
def kpool_seed_tail_kernel[dtype: DType, TailLayoutType: TensorLayout, tail_origin: MutOrigin, KLayoutType: TensorLayout, k_origin: ImmOrigin, GateLayoutType: TensorLayout, gate_origin: ImmOrigin, IROLayoutType: TensorLayout, iro_origin: ImmOrigin, CacheLenLayoutType: TensorLayout, SlotLayoutType: TensorLayout, slot_origin: ImmOrigin, TailEngine: TensorEngine, KEngine: TensorEngine, GateEngine: TensorEngine, IROEngine: TensorEngine, CacheLenEngine: TensorEngine, SlotEngine: TensorEngine, head_dim: Int, kpool: Int](tail: TileTensor[dtype, TailLayoutType, tail_origin, Engine=TailEngine], k: TileTensor[dtype, KLayoutType, k_origin, Engine=KEngine], gate: TileTensor[dtype, GateLayoutType, gate_origin, Engine=GateEngine], input_row_offsets: TileTensor[.uint32, IROLayoutType, iro_origin, Engine=IROEngine], cache_lengths: TileTensor[.uint32, CacheLenLayoutType, ImmutAnyOrigin, Engine=CacheLenEngine], slot_idx: TileTensor[.uint32, SlotLayoutType, slot_origin, Engine=SlotEngine], num_requests: Int32)
Stashes a prefill chunk's trailing tokens into the tail ring.
Compression writes only whole pools, so the tokens after the last complete pool have nowhere else to go. They are the members of the request's in-progress pool, and decode reads them back from the ring.
Only the trailing tokens this call owns are written, so a chunked prefill lands where a single one does.
Parameters:
- dtype (
DType): Element type oftail,kandgate. - TailLayoutType (
TensorLayout): Layout oftail. - tail_origin (
MutOrigin): Origin oftail. - KLayoutType (
TensorLayout): Layout ofk. - k_origin (
ImmOrigin): Origin ofk. - GateLayoutType (
TensorLayout): Layout ofgate. - gate_origin (
ImmOrigin): Origin ofgate. - IROLayoutType (
TensorLayout): Layout ofinput_row_offsets. - iro_origin (
ImmOrigin): Origin ofinput_row_offsets. - CacheLenLayoutType (
TensorLayout): Layout ofcache_lengths. - SlotLayoutType (
TensorLayout): Layout ofslot_idx. - slot_origin (
ImmOrigin): Origin ofslot_idx. - TailEngine (
TensorEngine): Engine oftail. - KEngine (
TensorEngine): Engine ofk. - GateEngine (
TensorEngine): Engine ofgate. - IROEngine (
TensorEngine): Engine ofinput_row_offsets. - CacheLenEngine (
TensorEngine): Engine ofcache_lengths. - SlotEngine (
TensorEngine): Engine ofslot_idx. - head_dim (
Int): Channels per key; also the block width. - kpool (
Int): Tokens per pool.
Args:
- tail (
TileTensor[dtype, TailLayoutType, tail_origin, Engine=TailEngine]): Per-slot ring,[max_slots, 2, kpool, head_dim]. Index 0 holds keys, index 1 holds gate scores. - k (
TileTensor[dtype, KLayoutType, k_origin, Engine=KEngine]): Layer-normed indexer keys,[total_tokens, head_dim]. - gate (
TileTensor[dtype, GateLayoutType, gate_origin, Engine=GateEngine]): Per-token gate scores,[total_tokens, head_dim]. - input_row_offsets (
TileTensor[.uint32, IROLayoutType, iro_origin, Engine=IROEngine]): Token row offsets per request,[batch_size + 1]. - cache_lengths (
TileTensor[.uint32, CacheLenLayoutType, ImmutAnyOrigin, Engine=CacheLenEngine]): Cached-prefix length per request,[batch_size]. - slot_idx (
TileTensor[.uint32, SlotLayoutType, slot_origin, Engine=SlotEngine]): Ring slot owned by each batch row,[num_requests],uint32. - num_requests (
Int32): Requests actually present.