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

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 concrete KVCacheT type 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

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

static def get_type_name() -> String

Returns:

String

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:

Pointer[Scalar[Self.dtype], ImmutAnyOrigin]

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:

Pointer[Scalar[Self.scale_dtype], ImmutAnyOrigin]

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:

SIMD[Self.scale_dtype, width]

cache_length

def cache_length(self, batch_idx: Int) -> Int

Returns:

Int

max_context_length

def max_context_length(self) -> UInt32

Returns:

UInt32

num_kv_rows

def num_kv_rows(self) -> Int

Returns the total number of virtual rows in the KV memory view.

Returns:

Int

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:

UInt32

get_tma_row

def get_tma_row(self, encoded_index: Int32) -> Int32

Convert an encoded sparse index to a physical TMA row.

Returns:

Int32

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:

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]()]

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_rope_tma_tile

def create_rope_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int, padded_depth: Int](self, ctx: DeviceContext, out tma: TMATensorTile[DType.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:

TMATensorTile[DType.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]()]

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) in tma_dtype elements.
  • tile_stride (Int): Row stride in elements in global memory. Defaults to tile_width.
  • 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.
  • tma_dtype (DType): Data type used for the TMA descriptor. Defaults to Self.dtype.
  • 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))]

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[DType.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:

TMATensorTile[DType.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))]

scales_raw_ptr

def scales_raw_ptr(self) -> Pointer[Float32, MutAnyOrigin]

Returns a dangling pointer. KVCacheScalesMHAOperand already points to the scales pointer.

Returns:

Pointer[Float32, MutAnyOrigin]