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 function

load_AB

def load_AB[a_type: DType, b_type: DType, a_scales_type: DType, a_rank: Int, a_tile_shape: IndexList[a_rank], a_desc_shape: IndexList[a_rank], b_rank: Int, b_tile_shape: IndexList[b_rank], b_desc_shape: IndexList[b_rank], a_scales_rank: Int, a_scales_tile_shape: IndexList[a_scales_rank], a_scales_desc_shape: IndexList[a_scales_rank], a_dim0: Int, a_dim1: Int, a_num_tiles: Int, a_swizzle_bytes: Int, b_dim0: Int, b_dim1: Int, b_num_tiles: Int, b_swizzle_bytes: Int, a_scales_dim0: Int, a_scales_dim1: Int, a_scales_num_tiles: Int, num_pipeline_stages: Int, /, *, block_tile_shape: IndexList[Int(3)], mma_shape: IndexList[Int(3)], cta_group: Int = Int(1)](a_tma_op: TMATensorTile[a_type, a_rank, a_tile_shape, a_desc_shape], b_tma_op: TMATensorTile[b_type, b_rank, b_tile_shape, b_desc_shape], a_scales_tma_op: TMATensorTile[a_scales_type, a_scales_rank, a_scales_tile_shape, a_scales_desc_shape], a_smem_tiles: SMemTileArray2D[a_type, a_dim0, a_dim1, a_num_tiles, a_swizzle_bytes], b_smem_tiles: SMemTileArray2D[b_type, b_dim0, b_dim1, b_num_tiles, b_swizzle_bytes], a_scales_smem_tiles: SMemTileArray2DRowMajor[a_scales_type, a_scales_dim0, a_scales_dim1, a_scales_num_tiles], load_mma_pipeline: ProducerConsumerPipeline[num_pipeline_stages], peer_cta_coord: Tuple[Int, Int, Int], work_tile_coord: Tuple[Int, Int], a_multicast_mask: UInt16, b_multicast_mask: UInt16, iter_idx: Int, elect_one_cta: Bool)

Loads A, B, and A-scales tiles into shared memory via multicast TMA for one pipeline stage.

Parameters:

  • ​a_type (DType): FP8 element type of the A operand matrix.
  • ​b_type (DType): FP8 element type of the B operand matrix.
  • ​a_scales_type (DType): Element type of the A blockwise scales (float32).
  • ​a_rank (Int): Number of dimensions in the A TMA descriptor.
  • ​a_tile_shape (IndexList[a_rank]): Shared-memory tile shape for each A TMA load.
  • ​a_desc_shape (IndexList[a_rank]): Copy box shape governing bytes per A TMA load.
  • ​b_rank (Int): Number of dimensions in the B TMA descriptor.
  • ​b_tile_shape (IndexList[b_rank]): Shared-memory tile shape for each B TMA load.
  • ​b_desc_shape (IndexList[b_rank]): Copy box shape governing bytes per B TMA load.
  • ​a_scales_rank (Int): Number of dimensions in the A scales TMA descriptor.
  • ​a_scales_tile_shape (IndexList[a_scales_rank]): Shared-memory tile shape for each A scales TMA load.
  • ​a_scales_desc_shape (IndexList[a_scales_rank]): Copy box shape governing bytes per A scales TMA load.
  • ​a_dim0 (Int): Row extent of each A shared-memory tile.
  • ​a_dim1 (Int): Column extent of each A shared-memory tile.
  • ​a_num_tiles (Int): Number of A tiles in the shared-memory tile array.
  • ​a_swizzle_bytes (Int): Swizzle stride in bytes for the A shared-memory tile.
  • ​b_dim0 (Int): Row extent of each B shared-memory tile.
  • ​b_dim1 (Int): Column extent of each B shared-memory tile.
  • ​b_num_tiles (Int): Number of B tiles in the shared-memory tile array.
  • ​b_swizzle_bytes (Int): Swizzle stride in bytes for the B shared-memory tile.
  • ​a_scales_dim0 (Int): Row extent of each A scales shared-memory tile.
  • ​a_scales_dim1 (Int): Column extent of each A scales shared-memory tile.
  • ​a_scales_num_tiles (Int): Number of A scales tiles in the shared-memory tile array.
  • ​num_pipeline_stages (Int): Number of double-buffered stages for A, B, and A scales TMA loads.
  • ​block_tile_shape (IndexList[Int(3)]): GEMM block tile shape (BM, BN, BK).
  • ​mma_shape (IndexList[Int(3)]): Tensor-core MMA shape (MMA_M, MMA_N, MMA_K).
  • ​cta_group (Int): Number of CTAs cooperating per MMA group (defaults to 1).

Args:

Was this page helpful?