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 trait
MHAOperand
This serves as the trait to support arguments to our MHA kernel.
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
dtype
comptime dtype
Engine
comptime Engine
page_size
comptime page_size
quantization_enabled
comptime quantization_enabled = False
quantization_granularity
comptime quantization_granularity
scale_dtype
comptime scale_dtype
Required methods
block_paged_storage
def block_paged_storage[tile_size: Int](self, batch_idx: UInt32, start_tok_idx: UInt32, head_idx: UInt32, head_dim_idx: UInt32 = UInt32(0)) -> Self.Engine.StorageType[False, ImmUnsafeAnyOrigin, Self.dtype, ImmutAnyOrigin, AddressSpace.GENERIC]
Returns:
_Self.Engine.StorageType[False, ImmUnsafeAnyOrigin, _Self.dtype, ImmutAnyOrigin, AddressSpace.GENERIC]
block_paged_ptr
def block_paged_ptr[tile_size: Int](self, batch_idx: UInt32, start_tok_idx: UInt32, head_idx: UInt32, head_dim_idx: UInt32 = UInt32(0)) -> Pointer[Scalar[Self.dtype], ImmutAnyOrigin]
Returns:
scales_block_paged_ptr
def scales_block_paged_ptr(self, batch_idx: Int, start_tok_idx: Int, head_idx: Int, head_dim_idx: Int = Int(0)) -> Pointer[Scalar[Self.scale_dtype], ImmutAnyOrigin]
Returns:
load_scale
def load_scale[width: Int](self, batch_idx: Int, start_tok_idx: Int, head_idx: Int, head_dim_idx: Int) -> SIMD[Self.scale_dtype, width]
Returns:
cache_length
def cache_length(self, batch_idx: Int) -> Int
Returns the length of the cache for a given batch index.
Returns:
max_context_length
def max_context_length(self) -> UInt32
Returns the maximum cache length in a given batch index.
Returns:
num_kv_rows
def num_kv_rows(self) -> Int
Returns the total number of virtual rows in the KV memory view.
For paged caches this accounts for the paging stride so that TMA descriptors can be sized to cover the entire address space.
Returns:
row_idx
def row_idx(self, batch_idx: UInt32, start_tok_idx: UInt32) -> UInt32
Returns the row idx when viewing the memory as a matrix.
Returns:
get_tma_row
def get_tma_row(self, encoded_index: Int32) -> Int32
Convert an encoded sparse index to a physical TMA row.
For paged caches the encoded index is
physical_block * page_size + offset and this method returns
physical_block * stride + offset. Non-paged operands return
the encoded index unchanged.
Returns:
create_tma_tile
def create_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, depth: Int, BK: Int = padded_depth[_Self.dtype, swizzle_mode, depth](), fold_chunks: Int = Int(1), row_major: Bool = False](self, ctx: DeviceContext) -> TMATensorTile[Self.dtype, Int(3), _padded_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()]
Creates a TMA tile for efficient GPU memory transfers. This is useful for k-major MMA operations where we don't need to mask any extra rows.
When fold_chunks >= 2 the contiguous depth chunks are folded into one
rank-4 TMA descriptor (SM100 K-only optimization). The caller must pass the
value from kv_tma_fold_chunks and use the same value at the tma_copy_k
issue site. 1 (default) is the original per-chunk behavior.
When row_major is True (and fold_chunks >= 2) the fold uses the
rank-5 chunk-inner (page-dense) box instead of the rank-4 chunk-outer
box, so a tile can span multiple pages with one TMA per page. The same
value must be used at the tma_copy_k/tma_copy_v issue site and the
P@V MMA consumer descriptor.
Returns:
create_scale_tma_tile
def create_scale_tma_tile[BMN: Int](self, ctx: DeviceContext) -> TMATensorTile[Self.scale_dtype, Int(2), Index[Int, Int](Int(1), BMN)]
Creates a TMA tile for efficient GPU memory transfers. This is useful for m-major MMA operations where we don't need to mask any extra rows.
Returns:
TMATensorTile[_Self.scale_dtype, Int(2), Index[Int, Int](Int(1), BMN)]
create_rope_tma_tile
def create_rope_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int, padded_depth: Int](self, ctx: DeviceContext) -> TMATensorTile[.bfloat16, Int(3), _padded_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()]
Creates a BF16 TMA tile for the rope portion of the per-tensor rope-aware KV cache.
Returns:
create_gather4_tma_tile
def create_gather4_tma_tile[tile_width: Int, tile_stride: Int = tile_width, swizzle_mode: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_NONE, tile_height: Int = Int(4), tma_dtype: DType = _Self.dtype, l2_promotion: TensorMapL2Promotion = TensorMapL2Promotion.NONE](self, ctx: DeviceContext) -> TMATensorTile[tma_dtype, Int(2), IndexList(tile_height, _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]
Creates a 2D TMA gather4 descriptor for this operand.
The descriptor views the data as a flat 2D matrix of
[num_kv_rows, tile_width] and is configured for gather4 operations
that load 4 non-contiguous rows per TMA instruction. The box width
is derived from the swizzle mode; for SWIZZLE_NONE it equals
tile_width.
When tma_dtype differs from Self.dtype, the underlying data
pointer is bitcast to tma_dtype at descriptor creation time.
Parameters:
- tile_width (
Int): Number of elements per row to load (box width) intma_dtypeelements. - tile_stride (
Int): Row stride in elements in global memory. Defaults totile_width. Use a larger value when the global row is wider than the portion to load. - swizzle_mode (
TensorMapSwizzle): TMA swizzle mode for shared memory access pattern. Defaults to SWIZZLE_NONE. - tile_height (
Int): Number of rows in the tile. Must be a multiple of 4. Defaults to 4 for backward compatibility. - tma_dtype (
DType): The data type used for the TMA descriptor. Defaults toSelf.dtype. When different, the pointer is bitcast. - l2_promotion (
TensorMapL2Promotion): L2 cache promotion hint for TMA loads. Defaults to NONE.
Args:
- ctx (
DeviceContext): The CUDA device context used to create the TMA descriptor.
Returns:
TMATensorTile[tma_dtype, Int(2), IndexList(tile_height, _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]: A TMATensorTile with box width derived from the swizzle mode.
create_rope_gather4_tma_tile
def create_rope_gather4_tma_tile[tile_width: Int, padded_depth: Int, swizzle_mode: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_NONE, tile_height: Int = Int(4), l2_promotion: TensorMapL2Promotion = TensorMapL2Promotion.NONE](self, ctx: DeviceContext) -> TMATensorTile[.bfloat16, Int(2), IndexList(tile_height, _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]
Creates a BF16 gather4 TMA descriptor for the rope portion of the KV cache.
For the per-tensor rope-aware layout each token row is
padded_depth FP8 bytes (content) followed by BF16 rope elements.
This method offsets the base pointer by padded_depth bytes,
reinterprets as BF16, and creates a gather4 TMA descriptor.
Parameters:
- tile_width (
Int): Number of BF16 elements per row in global memory. - padded_depth (
Int): Byte offset from row start to the rope data. - swizzle_mode (
TensorMapSwizzle): TMA swizzle mode for shared memory access pattern. - tile_height (
Int): Number of rows in the tile. Must be a multiple of 4. - l2_promotion (
TensorMapL2Promotion): L2 cache promotion hint for TMA loads. Defaults to NONE.
Args:
- ctx (
DeviceContext): The CUDA device context used to create the TMA descriptor.
Returns:
TMATensorTile[.bfloat16, Int(2), IndexList(tile_height, _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]: A BF16 TMATensorTile configured for gather4.
scales_raw_ptr
def scales_raw_ptr(self) -> Pointer[Float32, MutAnyOrigin]
Returns the base pointer to the quantization scales tensor.
Returns a null pointer for operands without quantization support.
Returns:
Provided methods
block_paged_tile
def block_paged_tile[layout_t: TensorLayout, //, tile_size: Int](self, batch_idx: UInt32, start_tok_idx: UInt32, head_idx: UInt32, layout_val: layout_t, head_dim_idx: UInt32 = UInt32(0)) -> TileTensor[Self.dtype, layout_t, ImmutAnyOrigin, Engine=Self.Engine]
Wraps block_paged_ptr in a TileTensor with the caller's layout.
Parameters:
- layout_t (
TensorLayout): TheTensorLayoutof the returnedTileTensor(inferred). - tile_size (
Int): Tile size in rows used to compute the paged block pointer.
Args:
- batch_idx (
UInt32): Batch index of the request. - start_tok_idx (
UInt32): Starting token index within the batch. - head_idx (
UInt32): KV head index. - layout_val (
layout_t): Concretelayout_tinstance describing the tile layout. - head_dim_idx (
UInt32): Index along the head dimension (defaults to 0).
Returns:
TileTensor[_Self.dtype, layout_t, ImmutAnyOrigin, Engine=_Self.Engine]
populate
def populate[BN: Int, base_alignment: Int, pair_cta: Bool = False, is_leader: Bool = True](self, batch_idx: UInt32, base_kv_row: UInt32) -> PagedRowIndices[BN, _Self.page_size, pair_cta, is_leader]
Populate a full PagedRowIndices[BN, ...] for a BN-row tile.
Returns the precomputed physical row indices for the num_pages
sub-tile pages covering the BN-row range starting at
base_kv_row for batch_idx. Both K's TMA (which may cover only
a subset in pair_cta mode) and V's TMA (which covers the full
range) can then consume the result without any lazy LUT lookup.
base_alignment is a comptime promise that
base_kv_row % base_alignment == 0 at runtime: typically
mask.start_column_alignment[...](). The PagedKVCache
override uses it to pick the largest legal SIMD chunk for its
LUT vector load and to skip the intra-page divmod when
base_alignment % page_size == 0.
Default implementation: scalar loop over num_pages calls to
row_idx. Overrides (e.g. PagedKVCache) replace this with a
single SIMD load from the underlying lookup table.
Returns:
PagedRowIndices[BN, _Self.page_size, pair_cta, is_leader]
create_index_scale_tma_tile
def create_index_scale_tma_tile[TILE: Int](self, ctx: DeviceContext) -> TMATensorTile[Self.dtype, Int(2), Index[Int, Int](Int(1), flat_scale_window[_Self.dtype, TILE]())]
Creates a flat TMA tile over the scale pool this operand addresses.
Only meaningful for an operand that IS a scale pool, so the default
rejects; block_paged_ptr must already address scales. The box is
TILE scales plus one alignment unit of slack -- a TMA's global start
must be 16-byte aligned and a key index is the innermost coordinate
here, so an unaligned caller rounds its base down and skips the
residual.
Returns:
TMATensorTile[_Self.dtype, Int(2), Index[Int, Int](Int(1), flat_scale_window[_Self.dtype, TILE]())]