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

MLAKPoolExpandTopK

struct MLAKPoolExpandTopK

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

Implemented traits​

AnyType, Deinitable, Movable

Methods​

execute​

static def execute[*, kpool: Int, pool_topk: Int, always_select_tail: Bool](out_indices: ManagedTensorSlice[IOSpec[_, _].Output, static_spec=out_indices.static_spec], pool_ids: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=pool_ids.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], ctx: DeviceContext)

Turns selected pool ids back into the token positions they cover.

The indexer selects pools; attention reads tokens. Pool p covers positions [p * kpool, (p + 1) * kpool), so the selection widens from pool_topk to pool_topk * kpool. An unselected slot expands to -1 in every one of its positions rather than to a clamped valid one, which would point attention at a token the indexer did not choose.

Parameters:

  • ​kpool (Int): Tokens per pool.
  • ​pool_topk (Int): Pools selected per query; the width of pool_ids.
  • ​always_select_tail (Bool): Append the kpool - 1 positions after the last complete pool -- the query's most recent tokens.

Args:

  • ​out_indices (ManagedTensorSlice[IOSpec[_, _].Output, static_spec=out_indices.static_spec]): Output [total_seq_len, pool_topk * kpool + tail] token positions, -1 where unused.
  • ​pool_ids (ManagedTensorSlice[IOSpec[_, _].Input, static_spec=pool_ids.static_spec]): Selected pool ids [total_seq_len, pool_topk].
  • ​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].
  • ​ctx (DeviceContext): Device context for GPU execution.

Was this page helpful?