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
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 ofpool_ids. - always_select_tail (
Bool): Append thekpool - 1positions 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,-1where 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.