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
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:
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.