IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo struct

LayoutTensorMHAOperand

struct LayoutTensorMHAOperand[origin: ImmOrigin, scale_origin: ImmOrigin, //, dtype_: DType, buffer_layout: TensorLayout, scale_dtype_: DType = DType.float32, scale_buffer_layout: TensorLayout = Layout[*?, *?]]

An implementation for contiguous tensor arguments to MHA kernels.

Parameters​

  • ​origin (ImmOrigin): Origin of the K/V buffer TileTensor (inferred).
  • ​scale_origin (ImmOrigin): Origin of the scale_buffer TileTensor (inferred).
  • ​dtype_ (DType): Element type of the K/V buffer.
  • ​buffer_layout (TensorLayout): TensorLayout of the K/V buffer.
  • ​scale_dtype_ (DType): Element type of the scale_buffer (defaults to DType.float32).
  • ​scale_buffer_layout (TensorLayout): TensorLayout of the scale_buffer. A rank-0 layout disables quantization (defaults to a rank-0 MixedLayout).

Fields​

  • ​buffer (TileTensor[LayoutTensorMHAOperand[dtype_, buffer_layout, scale_dtype_, scale_buffer_layout].dtype, buffer_layout, origin]):
  • ​scale_buffer (TileTensor[LayoutTensorMHAOperand[dtype_, buffer_layout, scale_dtype_, scale_buffer_layout].scale_dtype, scale_buffer_layout, scale_origin]):

Implemented traits​

AnyType, Copyable, DevicePassable, ImplicitlyCopyable, ImplicitlyDeletable, MHAOperand, Movable, RegisterPassable, TrivialRegisterPassable

comptime members​

device_type​

comptime device_type = LayoutTensorMHAOperand[dtype_, buffer_layout, scale_dtype_, scale_buffer_layout]

dtype​

comptime dtype = dtype_

layout_dim​

comptime layout_dim = buffer_layout.__shape_types[SIMDLength((buffer_layout.rank - Int(1)))].static_value

layout_rank​

comptime layout_rank = buffer_layout.rank

page_size​

comptime page_size = 0

quantization_enabled​

comptime quantization_enabled = (scale_buffer_layout.rank != Int(0))

quantization_granularity​

comptime quantization_granularity = ceildiv(buffer_layout.__shape_types[(add buffer_layout.rank, -1)].static_value, scale_buffer_layout.__shape_types[(add scale_buffer_layout.rank, -1)].static_value if (xor (eq scale_buffer_layout.rank, 0), True) else Int(1))

scale_dim​

comptime scale_dim = scale_buffer_layout.__shape_types[SIMDLength((scale_buffer_layout.rank - Int(1)))].static_value if (scale_buffer_layout.rank != Int(0)) else Int(1)

scale_dtype​

comptime scale_dtype = scale_dtype_

scale_rank​

comptime scale_rank = scale_buffer_layout.rank

Methods​

__init__​

def __init__(buffer: TileTensor[Self.dtype, buffer_layout, origin], scale_buffer: TileTensor[Self.scale_dtype, scale_buffer_layout, scale_origin] = _null_scale_tile_tensor[LayoutTensorMHAOperand[dtype_, buffer_layout, scale_dtype_, scale_buffer_layout].scale_dtype, scale_buffer_layout]()) -> 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, _safe=False]

Returns:

Pointer[Scalar[Self.dtype], ImmutAnyOrigin, _safe=False]

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, _safe=False]

Returns:

Pointer[Scalar[Self.scale_dtype], ImmutAnyOrigin, _safe=False]

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 (batch * seq_len).

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.

Non-paged operand: identity (no paging translation needed).

Returns:

Int32

create_tma_tile​

def create_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, depth: Int, BK: Int = padded_depth[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]()])

Creates a TMA tile for efficient GPU memory transfers.

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 LayoutTensorMHAOperand.

Parameters:

  • ​swizzle_mode (TensorMapSwizzle): TMA swizzle mode for shared memory access pattern.
  • ​BN (Int): Number of rows loaded per TMA instruction.
  • ​BK (Int): Block size along the depth dimension in BF16 elements.
  • ​padded_depth (Int): Byte offset from row start to the BF16 rope data, skipping the FP8 content prefix.

Args:

  • ​ctx (DeviceContext): The CUDA device context used to create the TMA descriptor.

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 = LayoutTensorMHAOperand[dtype_, buffer_layout, scale_dtype_, scale_buffer_layout].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))])

Creates a 2D TMA gather4 descriptor for this contiguous operand.

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. 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): Data type used for the TMA descriptor. Defaults to Self.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))]

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 LayoutTensorMHAOperand.

Parameters:

  • ​tile_width (Int): Number of BF16 elements per row in global memory.
  • ​padded_depth (Int): Byte offset from row start to the BF16 rope data, skipping the FP8 content prefix.
  • ​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.
  • ​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[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, _safe=False]

Returns a dangling pointer. Contiguous operands do not support quantization.

Returns:

Pointer[Float32, MutAnyOrigin, _safe=False]

Was this page helpful?