IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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​

AnyType, Deinitable, Movable

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:

  • ​head_dim (Int): Channels per key; also the kernel's block width.
  • ​kpool (Int): Tokens per pool.

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.

Was this page helpful?