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

MLAKPoolRingClose

struct MLAKPoolRingClose

Registers the mo.mla.kpool.ring_close 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], closed_pool: ManagedTensorSlice[IOSpec[_, _].Output, static_spec=closed_pool.static_spec], 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], ape: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=ape.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)

Closes each request's pending tail-ring pool, ragged and unconditional -- see kpool_ring_close_kernel's docstring for the pool-splitting semantics this implements.

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 [batch, head_dim] pooled keys, meaningful only where closed_pool is non-negative.
  • ​closed_pool (ManagedTensorSlice[IOSpec[_, _].Output, static_spec=closed_pool.static_spec]): Output [batch] pool id closed this call, or -1.
  • ​tail (ManagedTensorSlice[IOSpec[_, _].MutableInput, static_spec=tail.static_spec]): Persistent slot-indexed ring [max_slots, 2, kpool, head_dim]. Read only.
  • ​k (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec]): This call's layer-normed keys [total_tokens, head_dim].
  • ​gate (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gate.static_spec]): This call's 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 per request [batch + 1].
  • ​cache_lengths (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=cache_lengths.static_spec]): Cached-prefix length 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?