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
MLAKPoolCompress
struct MLAKPoolCompress
Registers the mo.mla.kpool.compress graph op with the graph compiler.
Implemented traits
Methods
execute
static def execute[*, head_dim: Int, kpool: Int](pooled: ManagedTensorSlice[IOSpec[_, _].Output, static_spec=pooled.static_spec], k: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec], gate: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gate.static_spec], ape: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=ape.static_spec], input_row_offsets: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input_row_offsets.static_spec], pool_row_offsets: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=pool_row_offsets.static_spec], cache_lengths: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=cache_lengths.static_spec], ctx: DeviceContext)
Compresses each complete k-pool into one candidate key.
Parameters:
Args:
- pooled (
ManagedTensorSlice[IOSpec[_, _].Output, static_spec=pooled.static_spec]): Output[total_pools, head_dim]pooled keys. - k (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec]): Layer-normed indexer keys[total_tokens, head_dim]. - gate (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gate.static_spec]): Per-token gate scores[total_tokens, head_dim]. - ape (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=ape.static_spec]): Within-pool position embedding[kpool, head_dim], float32. - input_row_offsets (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input_row_offsets.static_spec]): Token row offsets[batch + 1]. - pool_row_offsets (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=pool_row_offsets.static_spec]): Pool row offsets[batch + 1]. - cache_lengths (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=cache_lengths.static_spec]): Cached-prefix length per request[batch]. A pool covers absolute positions, so this is what places the call's tokens on the pool grid. - ctx (
DeviceContext): Device context for GPU execution.