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
KVCacheScalesMHAOperand
struct KVCacheScalesMHAOperand[cache_t: KVCacheT]
An MHAOperand that accesses the scales field of a KVCache.
This is useful for MLA attention where k_s (per-token scales) are stored in the scales field of the k cache with quantization_granularity = head_size. The scales have shape [num_blocks, page_size, num_heads, head_dim_granularity].
Parameters
- cache_t (
KVCacheT): The concreteKVCacheTtype whose scales field this operand accesses.
Fields
- cache (
cache_t):
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
MHAOperand,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
device_type
comptime device_type = KVCacheScalesMHAOperand[cache_t]
dtype
comptime dtype = cache_t.scale_dtype
Engine
comptime Engine = DefaultEngine
page_size
comptime page_size = cache_t.page_size_
quantization_enabled
comptime quantization_enabled = cache_t.quantization_enabled
quantization_granularity
comptime quantization_granularity = cache_t.quantization_granularity
scale_dtype
comptime scale_dtype = KVCacheScalesMHAOperand[cache_t].dtype
Methods
__init__
def __init__(cache: cache_t) -> Self
get_type_name
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:
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)) -> 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
max_context_length
num_kv_rows
def num_kv_rows(self) -> Int
Returns the total number of virtual rows in the KV memory view.
Returns:
row_idx
def row_idx(self, batch_idx: UInt32, start_tok_idx: UInt32) -> UInt32
Returns the row idx in the SCALE pool -- what this operand addresses.
The scale pool has its own lookup table and its own block stride, which
coincide with the value pool's only when the caller leaves the scales
LUT unset. block_paged_ptr already forwards to the scales, so this
must too.
Returns:
scale_tma_coords
def scale_tma_coords(self, batch_idx: UInt32, start_tok_idx: UInt32) -> Tuple[Int32, Int32]
The paged (row_in_block, block) coordinate; see the cache's own.
Args:
- batch_idx (
UInt32): Batch entry to address. - start_tok_idx (
UInt32): First token of the tile, within the entry.
Returns:
Tuple[Int32, Int32]: The row within the block, and the block.
get_tma_row
def get_tma_row(self, encoded_index: Int32) -> Int32
Convert an encoded sparse index to a physical TMA row.
Returns:
create_tma_tile
def create_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, depth: Int, BK: Int = padded_depth[cache_t.scale_dtype, swizzle_mode, depth](), fold_chunks: Int = Int(1), row_major: Bool = False](self, ctx: DeviceContext, out tma: 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]()])
TMA not supported for KVCacheScalesMHAOperand.
Returns:
create_scale_tma_tile
def create_scale_tma_tile[BMN: Int](self, ctx: DeviceContext, out tma: TMATensorTile[Self.scale_dtype, Int(2), Index[Int, Int](Int(1), BMN)])
Returns:
TMATensorTile[Self.scale_dtype, Int(2), Index[Int, Int](Int(1), BMN)]
create_index_scale_tma_tile
def create_index_scale_tma_tile[TILE: Int](self, ctx: DeviceContext, out tma: TMATensorTile[Self.dtype, Int(2), Index[Int, Int](Int(1), flat_scale_window[cache_t.scale_dtype, TILE]())])
A flat window on the cache's scale pool.
This operand IS the scales, so unlike create_scale_tma_tile -- which
describes a companion buffer and stays unimplemented here -- this one
is the operand's whole point.
Returns:
create_rope_tma_tile
def create_rope_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int, padded_depth: Int](self, ctx: DeviceContext, out tma: 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]()])
Not supported for KVCacheScalesMHAOperand.
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 = KVCacheScalesMHAOperand[cache_t].dtype, l2_promotion: TensorMapL2Promotion = TensorMapL2Promotion.NONE](self, ctx: DeviceContext, out tma: 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))])
Not supported for KVCacheScalesMHAOperand.
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. - swizzle_mode (
TensorMapSwizzle): TMA swizzle mode for shared memory access pattern. Defaults toSWIZZLE_NONE. - tile_height (
Int): Number of rows in the tile. Must be a multiple of 4. Defaults to 4. - tma_dtype (
DType): Data type used for the TMA descriptor. Defaults toSelf.dtype. - l2_promotion (
TensorMapL2Promotion): L2 cache promotion hint for TMA loads. Defaults toNONE.
Args:
- ctx (
DeviceContext): The CUDA device context used to create the TMA descriptor.
Returns:
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, out tma: 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))])
Not supported for KVCacheScalesMHAOperand.
Returns:
scales_raw_ptr
def scales_raw_ptr(self) -> Pointer[Float32, MutAnyOrigin]
Returns a dangling pointer. KVCacheScalesMHAOperand already points to the scales pointer.
Returns: