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_compress_kernel
def kpool_compress_kernel[dtype: DType, KLayoutType: TensorLayout, k_origin: ImmOrigin, GateLayoutType: TensorLayout, gate_origin: ImmOrigin, ApeLayoutType: TensorLayout, ape_origin: ImmOrigin, IROLayoutType: TensorLayout, iro_origin: ImmOrigin, PROLayoutType: TensorLayout, pro_origin: ImmOrigin, CacheLenLayoutType: TensorLayout, OutLayoutType: TensorLayout, out_origin: MutOrigin, head_dim: Int, kpool: Int](pooled: TileTensor[dtype, OutLayoutType, out_origin], k: TileTensor[dtype, KLayoutType, k_origin], gate: TileTensor[dtype, GateLayoutType, gate_origin], ape: TileTensor[.float32, ApeLayoutType, ape_origin], input_row_offsets: TileTensor[.uint32, IROLayoutType, iro_origin], pool_row_offsets: TileTensor[.uint32, PROLayoutType, pro_origin], cache_lengths: TileTensor[.uint32, CacheLenLayoutType, ImmutAnyOrigin])
Builds one pooled key per block; one thread per channel.
Parameters:
- dtype (
DType): Element type ofk,gateandpooled. - KLayoutType (
TensorLayout): Layout ofk. - k_origin (
ImmOrigin): Origin ofk. - GateLayoutType (
TensorLayout): Layout ofgate. - gate_origin (
ImmOrigin): Origin ofgate. - ApeLayoutType (
TensorLayout): Layout ofape. - ape_origin (
ImmOrigin): Origin ofape. - IROLayoutType (
TensorLayout): Layout ofinput_row_offsets. - iro_origin (
ImmOrigin): Origin ofinput_row_offsets. - PROLayoutType (
TensorLayout): Layout ofpool_row_offsets. - pro_origin (
ImmOrigin): Origin ofpool_row_offsets. - CacheLenLayoutType (
TensorLayout): Layout ofcache_lengths. - OutLayoutType (
TensorLayout): Layout ofpooled. - out_origin (
MutOrigin): Origin ofpooled. - head_dim (
Int): Channels per key; also the block width. - kpool (
Int): Tokens per pool.
Args:
- pooled (
TileTensor[dtype, OutLayoutType, out_origin]): Output[total_pools, head_dim], wheretotal_poolsis the last entry ofpool_row_offsets. - k (
TileTensor[dtype, KLayoutType, k_origin]): Layer-normed indexer keys,[total_tokens, head_dim]. - gate (
TileTensor[dtype, GateLayoutType, gate_origin]): Per-token gate scores,[total_tokens, head_dim]. - ape (
TileTensor[.float32, ApeLayoutType, ape_origin]): Within-pool position embedding,[kpool, head_dim], f32. - input_row_offsets (
TileTensor[.uint32, IROLayoutType, iro_origin]): Token row offsets per request,[batch_size + 1]. - pool_row_offsets (
TileTensor[.uint32, PROLayoutType, pro_origin]): Output row offsets per request,[batch_size + 1]. Requestbowns output rows[pool_row_offsets[b], pool_row_offsets[b + 1]), one per pool built here. - cache_lengths (
TileTensor[.uint32, CacheLenLayoutType, ImmutAnyOrigin]): Cached-prefix length per request,[batch_size]. A pool covers absolute positions, so this is what places the call's tokens on the pool grid.