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

consumer_main_loop

def consumer_main_loop[accum_type: DType, c_type: DType, a_type: DType, b_type: DType, sfa_dtype: DType, sfb_dtype: DType, 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_swizzle: TensorMapSwizzle, b_swizzle: TensorMapSwizzle, transpose_b: Bool, pipeline_stages: Int, scaling_kind: UMMAKind, num_group_pipeline_stages: Int, /, *, block_tile_shape: IndexList[Int(3)], mma_shape: IndexList[Int(3)], SFA_NUM_COLS: Int, SFB_NUM_COLS: Int, cta_group: Int = Int(1), cluster_shape: IndexList[Int(3)] = Index[Int, Int, Int](Int(1), Int(1), Int(1)), k_group_size: Int = Int(1)](tmem_addr: UInt32, sfa_tmem: UInt32, sfb_tmem: UInt32, 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], sfa_smem_tiles: SMemTileArrayWithLayout[sfa_dtype], sfb_smem_tiles: SMemTileArrayWithLayout[sfb_dtype], load_mma_pipeline: ProducerConsumerPipeline[pipeline_stages], sfb_pipeline: ProducerConsumerPipeline[num_group_pipeline_stages], sfb_ready_mbars: Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False], mut sfb_ready_state: PipelineState[num_group_pipeline_stages], mma_op: MmaOpSM100_BlockScaled_SS[c_type, a_type, b_type, sfa_dtype, sfb_dtype, scaling_kind, block_tile_shape, mma_shape, accum_type=accum_type, cta_group=cta_group, cluster_shape=cluster_shape, a_swizzle=a_swizzle, b_swizzle=b_swizzle, transpose_b=transpose_b, enable_small_sfb=True], elect_one_warp: Bool, iter_idx: UInt32, k_start: UInt32, work_tile_coord: Tuple[Int, Int])

TileTensor overload of consumer_main_loop for block-scaled MMA (small BN).

Accepts SMemTileArray2D for A/B tiles and SMemTileArrayWithLayout for scale factor tiles, calling the TileTensor MMA overload.

Parameters:

  • ​accum_type (DType): Element dtype of the MMA accumulator.
  • ​c_type (DType): Element dtype of the C output matrix.
  • ​a_type (DType): Element dtype of the A operand matrix.
  • ​b_type (DType): Element dtype of the B operand matrix.
  • ​sfa_dtype (DType): Element dtype of the A scale factors.
  • ​sfb_dtype (DType): Element dtype of the B scale factors.
  • ​a_dim0 (Int): Row count of each A SMEM tile.
  • ​a_dim1 (Int): Column count of each A SMEM tile.
  • ​a_num_tiles (Int): Total number of A SMEM tiles across all pipeline stages.
  • ​a_swizzle_bytes (Int): Swizzle stride in bytes for the A SMEM tiles.
  • ​b_dim0 (Int): Row count of each B SMEM tile.
  • ​b_dim1 (Int): Column count of each B SMEM tile.
  • ​b_num_tiles (Int): Total number of B SMEM tiles across all pipeline stages.
  • ​b_swizzle_bytes (Int): Swizzle stride in bytes for the B SMEM tiles.
  • ​a_swizzle (TensorMapSwizzle): TMA swizzle mode for the A operand.
  • ​b_swizzle (TensorMapSwizzle): TMA swizzle mode for the B operand.
  • ​transpose_b (Bool): Whether B is stored in K-major (transposed) layout.
  • ​pipeline_stages (Int): Number of producer/consumer stages in the A/B/SFA load and MMA pipeline.
  • ​scaling_kind (UMMAKind): Block-scaled MMA kind (for example MXFP8 or NVFP4).
  • ​num_group_pipeline_stages (Int): Number of producer/consumer stages in the SFB load pipeline, counted in K-group units.
  • ​block_tile_shape (IndexList[Int(3)]): Block tile shape as (BM, BN, BK) in elements.
  • ​mma_shape (IndexList[Int(3)]): MMA atom shape as (MMA_M, MMA_N, MMA_K) in elements.
  • ​SFA_NUM_COLS (Int): Number of TMEM columns occupied per pipeline stage by the A scale factors.
  • ​SFB_NUM_COLS (Int): Number of TMEM columns occupied per pipeline stage by the B scale factors.
  • ​cta_group (Int): Number of CTAs cooperating per MMA group (defaults to 1).
  • ​cluster_shape (IndexList[Int(3)]): Cluster dimensions in (M, N, K) thread blocks (defaults to (1, 1, 1)).
  • ​k_group_size (Int): Number of K-tiles loaded per pipeline stage (defaults to 1).

Args:

Was this page helpful?