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_tile_rank: Int, a_tile_shape: IndexList[a_tile_rank], a_desc_shape: IndexList[a_tile_rank], b_tile_rank: Int, b_tile_shape: IndexList[b_tile_rank], b_desc_shape: IndexList[b_tile_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, num_pipeline_stages: Int, //, *, block_tile_shape: IndexList[Int(3)], mma_shape: IndexList[Int(3)], cta_group: Int = Int(1), a_plane_splits: IndexList[Int(2)] = Index[Int, Int](Int(0), Int(0))](expert_ids: Pointer[Int32, _safe=False], a_tma_op: TMATensorTile[a_type, a_tile_rank, a_tile_shape, a_desc_shape], b_tma_op: TMATensorTile[b_type, b_tile_rank, b_tile_shape, b_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], mma_mbar: Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False], tma_mbar: Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False], producer_phase: PipelineState[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: UInt32, elect_one_cta: Bool, scheduler: TileScheduler[static_MN=scheduler.static_MN, tile_shape=scheduler.tile_shape, cluster=scheduler.cluster, cta_group=scheduler.cta_group, swizzle=scheduler.swizzle, swapAB=scheduler.swapAB], qkv_plane_stride: Int = Int(0))

Loads A and B tiles from global memory into shared memory via TMA multicast.

Issues asynchronous multicast TMA loads for the current pipeline stage, addressing the expert-local slice of A and the shared B slice, then signals completion through the TMA mbarrier.

Parameters:

  • ​a_type (DType): Element type of the A operand tiles (inferred).
  • ​b_type (DType): Element type of the B operand tiles (inferred).
  • ​a_tile_rank (Int): Rank of the A TMA tile descriptor (inferred).
  • ​a_tile_shape (IndexList[a_tile_rank]): Element shape of one A TMA tile (inferred).
  • ​a_desc_shape (IndexList[a_tile_rank]): Descriptor shape of the A TMA tile, used to compute per-load element counts and row stride (inferred).
  • ​b_tile_rank (Int): Rank of the B TMA tile descriptor (inferred).
  • ​b_tile_shape (IndexList[b_tile_rank]): Element shape of one B TMA tile (inferred).
  • ​b_desc_shape (IndexList[b_tile_rank]): Descriptor shape of the B TMA tile, used to compute per-load element counts and row stride (inferred).
  • ​a_dim0 (Int): Number of rows in one A shared-memory tile (inferred).
  • ​a_dim1 (Int): Number of columns in one A shared-memory tile (inferred).
  • ​a_num_tiles (Int): Number of pipeline stages in the A shared-memory tile array (inferred).
  • ​a_swizzle_bytes (Int): Swizzle granularity in bytes for A shared-memory tiles (inferred).
  • ​b_dim0 (Int): Number of rows in one B shared-memory tile (inferred).
  • ​b_dim1 (Int): Number of columns in one B shared-memory tile (inferred).
  • ​b_num_tiles (Int): Number of pipeline stages in the B shared-memory tile array (inferred).
  • ​b_swizzle_bytes (Int): Swizzle granularity in bytes for B shared-memory tiles (inferred).
  • ​num_pipeline_stages (Int): Number of double-buffered pipeline stages for A and B shared-memory tiles (inferred).
  • ​block_tile_shape (IndexList[Int(3)]): Block tile shape [BM, BN, BK] partitioning the GEMM into work tiles.
  • ​mma_shape (IndexList[Int(3)]): MMA instruction shape [MMA_M, MMA_N, MMA_K].
  • ​cta_group (Int): Number of CTAs cooperating per MMA along the M dimension (defaults to 1).
  • ​a_plane_splits (IndexList[Int(2)]): Per-plane split sizes for fused LoRA QKV A-plane row offsetting; (0, 0) disables it (defaults to (0, 0)).

Args:

Was this page helpful?