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

MLAKPoolSeedTail

struct MLAKPoolSeedTail

Registers the mo.mla.kpool.seed_tail graph op with the graph compiler.

Implemented traits​

AnyType, Deinitable, Movable

Methods​

execute​

static def execute[*, head_dim: Int, kpool: Int](tail: ManagedTensorSlice[IOSpec[_, _].MutableInput, static_spec=tail.static_spec], k: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec], gate: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gate.static_spec], input_row_offsets: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input_row_offsets.static_spec], cache_lengths: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=cache_lengths.static_spec], slot_idx: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=slot_idx.static_spec], ctx: DeviceContext)

Stashes a prefill chunk's trailing tokens into the tail ring.

Whole pools compress directly from k/gate; the tokens after the last complete pool are that request's in-progress pool and have nowhere else to go until later tokens complete it, so they seed this per-request ring. mo.mla.kpool.ring_close reads them back and closes the pool once it fills.

Parameters:

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

Args:

  • ​tail (ManagedTensorSlice[IOSpec[_, _].MutableInput, static_spec=tail.static_spec]): Persistent slot-indexed ring [max_slots, 2, kpool, head_dim], mutated in place.
  • ​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].
  • ​input_row_offsets (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input_row_offsets.static_spec]): Token row offsets per request [batch + 1].
  • ​cache_lengths (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=cache_lengths.static_spec]): Cached tokens per request [batch].
  • ​slot_idx (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=slot_idx.static_spec]): Ring slot per request [batch].
  • ​ctx (DeviceContext): Device context for GPU execution.

Was this page helpful?