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 function

kpool_tail_update_kernel

def kpool_tail_update_kernel[dtype: DType, TailLayoutType: TensorLayout, tail_origin: MutOrigin, OutLayoutType: TensorLayout, out_origin: MutOrigin, ClosedLayoutType: TensorLayout, closed_origin: MutOrigin, KLayoutType: TensorLayout, k_origin: ImmOrigin, GateLayoutType: TensorLayout, gate_origin: ImmOrigin, ApeLayoutType: TensorLayout, ape_origin: ImmOrigin, PosLayoutType: TensorLayout, pos_origin: ImmOrigin, SlotLayoutType: TensorLayout, slot_origin: ImmOrigin, head_dim: Int, kpool: Int, next_n: Int = Int(1)](tail: TileTensor[dtype, TailLayoutType, tail_origin], pooled: TileTensor[dtype, OutLayoutType, out_origin], closed_pool: TileTensor[.int32, ClosedLayoutType, closed_origin], k: TileTensor[dtype, KLayoutType, k_origin], gate: TileTensor[dtype, GateLayoutType, gate_origin], ape: TileTensor[.float32, ApeLayoutType, ape_origin], positions: TileTensor[.int32, PosLayoutType, pos_origin], slot_idx: TileTensor[.uint32, SlotLayoutType, slot_origin], num_requests: Int32)

Stashes a request's new tokens, and closes each pool as it fills.

A decoded token cannot be pooled on arrival, because its pool-mates arrived on earlier steps and have left the batch. Each request keeps its in-progress pool in tail, a ring of kpool slots addressed by position % kpool.

A speculative step appends next_n tokens at once, so several pools can close in one call.

The ring is indexed by slot_idx[r], not by r. A batch reorders between steps, so row r is not always the same request.

Every real token stashes, whether or not it closes a pool.

Rejected speculative tokens are the caller's problem. The ring holds no pointer to rewind, so a rejected token that has already stashed stays.

Parameters:

  • ​dtype (DType): Element type of tail, k, gate and pooled.
  • ​TailLayoutType (TensorLayout): Layout of tail.
  • ​tail_origin (MutOrigin): Origin of tail.
  • ​OutLayoutType (TensorLayout): Layout of pooled.
  • ​out_origin (MutOrigin): Origin of pooled.
  • ​ClosedLayoutType (TensorLayout): Layout of closed_pool.
  • ​closed_origin (MutOrigin): Origin of closed_pool.
  • ​KLayoutType (TensorLayout): Layout of k.
  • ​k_origin (ImmOrigin): Origin of k.
  • ​GateLayoutType (TensorLayout): Layout of gate.
  • ​gate_origin (ImmOrigin): Origin of gate.
  • ​ApeLayoutType (TensorLayout): Layout of ape.
  • ​ape_origin (ImmOrigin): Origin of ape.
  • ​PosLayoutType (TensorLayout): Layout of positions.
  • ​pos_origin (ImmOrigin): Origin of positions.
  • ​SlotLayoutType (TensorLayout): Layout of slot_idx.
  • ​slot_origin (ImmOrigin): Origin of slot_idx.
  • ​head_dim (Int): Channels per key; also the block width.
  • ​kpool (Int): Tokens per pool.
  • ​next_n (Int): Tokens appended per request per call.

Args:

Was this page helpful?