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
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:
Args:
- pooled (
ManagedTensorSlice[IOSpec[_, _].Output, static_spec=pooled.static_spec]): Output[batch, head_dim]pooled keys, meaningful only whereclosed_poolis 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.