Skip to content
KernelIndex
Search⌘K

submission 409293

moshe_05859 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1091 lines, June 9 Researcher Reciprocity License v1.0.

cudeepy_nvfp4_group_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-409293?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
22.3µs
#122 of 310
2026-01-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:89cc877edaff14e84b91e7f9ed96ed443a061da7bc6bcd6f38c641cb9b8c255f
license declaredunknown
license concludedunknown
authorsmoshe_05859
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fused-epilogueself.epilog_warp_id = (0, 1, 2, 3) # 4 epilogue warps
mbarrierself.epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(self.epilog_warp_id))
persistent-kerneltile_sched_params: utils.PersistentTileSchedulerParams, group_count: cutlass.Constexpr,
shared-memoryreserved_smem_bytes = 1024
tcgen05self.cta_group = tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
tile-m = 128CLUSTER_TILE_M = 128 * CLUSTER_M
tile-n = 128CLUSTER_TILE_M, CLUSTER_TILE_N = 128 * CLUSTER_M, MMA_N * CLUSTER_N
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

cudeepy_nvfp4_group_gemm.py1091 lines
# SM100 Grouped BlockScaled GEMM Kernel (NVFP4) - Condensed Version
# Warp-specialized persistent kernel for Blackwell architecture

import argparse, functools
from typing import List, Type, Tuple, Union
from inspect import isclass
import torch
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.torch as cutlass_torch
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import from_dlpack
from cutlass.cutlass_dsl import Boolean, extract_mlir_values, new_from_mlir_values


class GroupedWorkTileInfo:
    """Work tile info for grouped GEMM with group search result."""
    def __init__(self, tile_idx: cute.Coord, is_valid_tile: Boolean, group_search_result):
        self._tile_idx = tile_idx
        self._is_valid_tile = Boolean(is_valid_tile)
        self.group_search_result = group_search_result

    def __extract_mlir_values__(self) -> list:
        values = extract_mlir_values(self.tile_idx)
        values.extend(extract_mlir_values(self.is_valid_tile))
        values.extend(extract_mlir_values(self.group_search_result))
        return values

    def __new_from_mlir_values__(self, values: list) -> "GroupedWorkTileInfo":
        if len(values) != 11:
            raise ValueError("Length of mlir values extracted is incorrect.")
        return GroupedWorkTileInfo(
            new_from_mlir_values(self._tile_idx, values[:3]),
            new_from_mlir_values(self._is_valid_tile, [values[3]]),
            new_from_mlir_values(self.group_search_result, values[4:11]),
        )

    @property
    def is_valid_tile(self) -> Boolean:
        return self._is_valid_tile

    @property
    def tile_idx(self) -> cute.Coord:
        return self._tile_idx


class Sm100GroupedBlockScaledGemmKernel:
    """SM100 grouped blockscaled GEMM kernel with TMA + warp specialization."""
    reserved_smem_bytes = 1024
    bytes_per_tensormap = 128
    num_tensormaps = 5
    tensor_memory_management_bytes = 12

    def __init__(self, sf_vec_size: int, mma_tiler_mn: Tuple[int, int], cluster_shape_mn: Tuple[int, int], occupancy: int = 1, max_ab_stages: int = 0):
        self.acc_dtype = cutlass.Float32
        self.sf_vec_size = sf_vec_size
        self.use_2cta_instrs = mma_tiler_mn[0] == 256
        self.cluster_shape_mn = cluster_shape_mn
        self.mma_tiler = (*mma_tiler_mn, 1)
        self.cta_group = tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
        self.tensormap_update_mode = utils.TensorMapUpdateMode.SMEM
        self.occupancy = occupancy
        self.max_ab_stages = max_ab_stages  # 0 = auto, >0 = limit stages
        self.epilog_warp_id = (0, 1, 2, 3)  # 4 epilogue warps
        self.mma_warp_id = 4
        self.tma_warp_id = 5
        self.threads_per_cta = 32 * len((self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id))
        self.epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(self.epilog_warp_id))
        self.tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=2, num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)))
        self.tensormap_ab_init_barrier = pipeline.NamedBarrier(barrier_id=3, num_threads=64)
        self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
        self.num_tmem_alloc_cols = 512

    def _setup_attributes(self):
        """Configure MMA shapes, cluster layouts, multicast masks, stages, and smem layouts."""
        self.mma_inst_shape_mn = (self.mma_tiler[0], self.mma_tiler[1])
        self.mma_inst_shape_mn_sfb = (
            self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
            cute.round_up(self.mma_inst_shape_mn[1], 128),
        )
        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
            self.sf_vec_size, self.cta_group, self.mma_inst_shape_mn,
        )
        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
            self.sf_vec_size, tcgen05.CtaGroup.ONE, self.mma_inst_shape_mn_sfb,
        )
        mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
        mma_inst_tile_k = 4
        self.mma_tiler = (self.mma_inst_shape_mn[0], self.mma_inst_shape_mn[1], mma_inst_shape_k * mma_inst_tile_k)
        self.mma_tiler_sfb = (self.mma_inst_shape_mn_sfb[0], self.mma_inst_shape_mn_sfb[1], mma_inst_shape_k * mma_inst_tile_k)
        self.cta_tile_shape_mnk = (self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), self.mma_tiler[1], self.mma_tiler[2])
        self.cluster_tile_shape_mnk = tuple(x * y for x, y in zip(self.cta_tile_shape_mnk, (*self.cluster_shape_mn, 1)))
        self.cluster_layout_vmnk = cute.tiled_divide(cute.make_layout((*self.cluster_shape_mn, 1)), (tiled_mma.thr_id.shape,))
        self.cluster_layout_sfb_vmnk = cute.tiled_divide(cute.make_layout((*self.cluster_shape_mn, 1)), (tiled_mma_sfb.thr_id.shape,))
        self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2])
        self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1])
        self.num_mcast_ctas_sfb = cute.size(self.cluster_layout_sfb_vmnk.shape[1])
        self.is_a_mcast = self.num_mcast_ctas_a > 1
        self.is_b_mcast = self.num_mcast_ctas_b > 1
        self.is_sfb_mcast = self.num_mcast_ctas_sfb > 1
        self.epi_tile = sm100_utils.compute_epilogue_tile_shape(self.cta_tile_shape_mnk, self.use_2cta_instrs, self.c_layout, self.c_dtype)
        self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
            tiled_mma, self.mma_tiler, self.a_dtype, self.b_dtype, self.epi_tile,
            self.c_dtype, self.c_layout, self.sf_dtype, self.sf_vec_size, self.smem_capacity, self.occupancy,
        )
        # Apply max_ab_stages limit if set
        if self.max_ab_stages > 0 and self.num_ab_stage > self.max_ab_stages:
            self.num_ab_stage = self.max_ab_stages
        self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(tiled_mma, self.mma_tiler, self.a_dtype, self.num_ab_stage)
        self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(tiled_mma, self.mma_tiler, self.b_dtype, self.num_ab_stage)
        self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(tiled_mma, self.mma_tiler, self.sf_vec_size, self.num_ab_stage)
        self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(tiled_mma, self.mma_tiler, self.sf_vec_size, self.num_ab_stage)
        self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(self.c_dtype, self.c_layout, self.epi_tile, self.num_c_stage)
        mbar_smem_bytes = self._get_mbar_smem_bytes(num_acc_stage=self.num_acc_stage, num_ab_stage=self.num_ab_stage, num_c_stage=self.num_c_stage)
        tensormap_smem_bytes = self.bytes_per_tensormap * self.num_tensormaps
        if mbar_smem_bytes + tensormap_smem_bytes + self.tensor_memory_management_bytes > self.reserved_smem_bytes:
            raise ValueError(f"smem consumption exceeds reserved {self.reserved_smem_bytes}")

    @cute.jit
    def __call__(self, initial_a: cute.Tensor, initial_b: cute.Tensor, initial_c: cute.Tensor,
                 initial_sfa: cute.Tensor, initial_sfb: cute.Tensor, group_count: cutlass.Constexpr[int],
                 problem_shape_mnkl: cute.Tensor, strides_abc: cute.Tensor, tensor_address_abc: cute.Tensor,
                 tensor_address_sfasfb: cute.Tensor, total_num_clusters: cutlass.Constexpr[int],
                 tensormap_cute_tensor: cute.Tensor, max_active_clusters: cutlass.Constexpr[int]):
        """Setup and launch grouped GEMM kernel."""
        self.a_dtype, self.b_dtype = initial_a.element_type, initial_b.element_type
        self.sf_dtype, self.c_dtype = initial_sfa.element_type, initial_c.element_type
        self.is_nvfp4_output = self.c_dtype is cutlass.Float4E2M1FN
        self.a_major_mode = utils.LayoutEnum.from_tensor(initial_a).mma_major_mode()
        self.b_major_mode = utils.LayoutEnum.from_tensor(initial_b).mma_major_mode()
        self.c_layout = utils.LayoutEnum.from_tensor(initial_c)
        if cutlass.const_expr(self.a_dtype != self.b_dtype):
            raise TypeError(f"Type mismatch: {self.a_dtype} != {self.b_dtype}")
        self._setup_attributes()

        sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_a.shape, self.sf_vec_size)
        initial_sfa = cute.make_tensor(initial_sfa.iterator, sfa_layout)
        sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_b.shape, self.sf_vec_size)
        initial_sfb = cute.make_tensor(initial_sfb.iterator, sfb_layout)

        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
            self.sf_vec_size, self.cta_group, self.mma_inst_shape_mn,
        )
        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
            self.sf_vec_size, tcgen05.CtaGroup.ONE, self.mma_inst_shape_mn_sfb,
        )
        atom_thr_size = cute.size(tiled_mma.thr_id.shape)

        # TMA setup for A/B/SFA/SFB/C
        a_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
        a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
        tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(a_op, initial_a, a_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape)

        b_op = sm100_utils.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma.thr_id)
        b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
        tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(b_op, initial_b, b_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape)

        sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
        sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0))
        tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(sfa_op, initial_sfa, sfa_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape, internal_type=cutlass.Int16)

        sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id)
        sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0))
        tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(sfb_op, initial_sfb, sfb_smem_layout, self.mma_tiler_sfb, tiled_mma_sfb, self.cluster_layout_sfb_vmnk.shape, internal_type=cutlass.Int16)

        self.num_tma_load_bytes = (cute.size_in_bytes(self.a_dtype, a_smem_layout) + cute.size_in_bytes(self.b_dtype, b_smem_layout) +
                                   cute.size_in_bytes(self.sf_dtype, sfa_smem_layout) + cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)) * atom_thr_size

        epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
        tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(cpasync.CopyBulkTensorTileS2GOp(), initial_c, epi_smem_layout, self.epi_tile)

        self.tile_sched_params, grid = self._compute_grid(total_num_clusters, self.cluster_shape_mn, max_active_clusters)
        self.buffer_align_bytes = 1024
        self.size_tensormap_in_i64 = self.num_tensormaps * self.bytes_per_tensormap // 8

        @cute.struct
        class SharedStorage:
            tensormap_buffer: cute.struct.MemRange[cutlass.Int64, self.size_tensormap_in_i64]
            ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            tmem_dealloc_mbar_ptr: cutlass.Int64
            tmem_holding_buf: cutlass.Int32
            sC: cute.struct.Align[cute.struct.MemRange[self.c_dtype, cute.cosize(self.c_smem_layout_staged.outer)], self.buffer_align_bytes]
            sA: cute.struct.Align[cute.struct.MemRange[self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)], self.buffer_align_bytes]
            sB: cute.struct.Align[cute.struct.MemRange[self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)], self.buffer_align_bytes]
            sSFA: cute.struct.Align[cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)], self.buffer_align_bytes]
            sSFB: cute.struct.Align[cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)], self.buffer_align_bytes]

        self.shared_storage = SharedStorage
        self.kernel(tiled_mma, tiled_mma_sfb, tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b,
                    tma_atom_sfa, tma_tensor_sfa, tma_atom_sfb, tma_tensor_sfb, tma_atom_c, tma_tensor_c,
                    self.cluster_layout_vmnk, self.cluster_layout_sfb_vmnk,
                    self.a_smem_layout_staged, self.b_smem_layout_staged, self.sfa_smem_layout_staged,
                    self.sfb_smem_layout_staged, self.c_smem_layout_staged, self.epi_tile,
                    self.tile_sched_params, group_count, problem_shape_mnkl, strides_abc,
                    tensor_address_abc, tensor_address_sfasfb, tensormap_cute_tensor,
        ).launch(grid=grid, block=[self.threads_per_cta, 1, 1], cluster=(*self.cluster_shape_mn, 1),
                 smem=self.shared_storage.size_in_bytes(), min_blocks_per_mp=1)
        return

    @cute.kernel
    def kernel(self, tiled_mma: cute.TiledMma, tiled_mma_sfb: cute.TiledMma,
               tma_atom_a: cute.CopyAtom, mA_mkl: cute.Tensor, tma_atom_b: cute.CopyAtom, mB_nkl: cute.Tensor,
               tma_atom_sfa: cute.CopyAtom, mSFA_mkl: cute.Tensor, tma_atom_sfb: cute.CopyAtom, mSFB_nkl: cute.Tensor,
               tma_atom_c: cute.CopyAtom, mC_mnl: cute.Tensor,
               cluster_layout_vmnk: cute.Layout, cluster_layout_sfb_vmnk: cute.Layout,
               a_smem_layout_staged: cute.ComposedLayout, b_smem_layout_staged: cute.ComposedLayout,
               sfa_smem_layout_staged: cute.Layout, sfb_smem_layout_staged: cute.Layout,
               c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout], epi_tile: cute.Tile,
               tile_sched_params: utils.PersistentTileSchedulerParams, group_count: cutlass.Constexpr,
               problem_sizes_mnkl: cute.Tensor, strides_abc: cute.Tensor, ptrs_abc: cute.Tensor,
               ptrs_sfasfb: cute.Tensor, tensormaps: cute.Tensor):
        """GPU device kernel - warp-specialized with TMA/MMA/Epilogue warps."""
        warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx())
        if warp_idx == self.tma_warp_id:
            cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_a)
            cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_b)
            cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_sfa)
            cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_sfb)
            cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_c)

        use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
        bidx, bidy, bidz = cute.arch.block_idx()
        mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
        is_leader_cta = mma_tile_coord_v == 0
        cta_rank_in_cluster = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster())
        block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(cta_rank_in_cluster)
        block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(cta_rank_in_cluster)
        tidx, _, _ = cute.arch.thread_idx()

        # Allocate and init shared memory
        smem = utils.SmemAllocator()
        storage = smem.allocate(self.shared_storage)
        tensormap_smem_ptr = storage.tensormap_buffer.data_ptr()
        tensormap_a_smem_ptr = tensormap_smem_ptr
        tensormap_b_smem_ptr = tensormap_a_smem_ptr + self.bytes_per_tensormap // 8
        tensormap_sfa_smem_ptr = tensormap_b_smem_ptr + self.bytes_per_tensormap // 8
        tensormap_sfb_smem_ptr = tensormap_sfa_smem_ptr + self.bytes_per_tensormap // 8
        tensormap_c_smem_ptr = tensormap_sfb_smem_ptr + self.bytes_per_tensormap // 8
        tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
        tmem_holding_buf = storage.tmem_holding_buf

        # Init pipelines
        ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
        ab_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, num_tma_producer)
        ab_pipeline = pipeline.PipelineTmaUmma.create(
            barrier_storage=storage.ab_full_mbar_ptr.data_ptr(), num_stages=self.num_ab_stage,
            producer_group=ab_pipeline_producer_group, consumer_group=ab_pipeline_consumer_group,
            tx_count=self.num_tma_load_bytes, cta_layout_vmnk=cluster_layout_vmnk,
        )
        acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_acc_consumer_threads = len(self.epilog_warp_id) * (2 if use_2cta_instrs else 1)
        acc_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, num_acc_consumer_threads)
        acc_pipeline = pipeline.PipelineUmmaAsync.create(
            barrier_storage=storage.acc_full_mbar_ptr.data_ptr(), num_stages=self.num_acc_stage,
            producer_group=acc_pipeline_producer_group, consumer_group=acc_pipeline_consumer_group,
            cta_layout_vmnk=cluster_layout_vmnk,
        )
        if use_2cta_instrs:
            if warp_idx == self.tma_warp_id:
                with cute.arch.elect_one():
                    cute.arch.mbarrier_init(tmem_dealloc_mbar_ptr, 32)

        pipeline_init_arrive(cluster_shape_mn=self.cluster_shape_mn, is_relaxed=True)

        # Setup smem tensors
        sC = storage.sC.get_tensor(c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner)
        sA = storage.sA.get_tensor(a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner)
        sB = storage.sB.get_tensor(b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner)
        sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
        sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)

        # Multicast masks
        a_full_mcast_mask, b_full_mcast_mask, sfa_full_mcast_mask, sfb_full_mcast_mask = None, None, None, None
        if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
            a_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2)
            b_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1)
            sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2)
            sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1)

        # Local_tile partition global tensors
        gA_mkl = cute.local_tile(mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None))
        gB_nkl = cute.local_tile(mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None))
        gSFA_mkl = cute.local_tile(mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None))
        gSFB_nkl = cute.local_tile(mSFB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None))
        gC_mnl = cute.local_tile(mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None))

        # Partition for TiledMMA
        thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
        thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
        tCgA = thr_mma.partition_A(gA_mkl)
        tCgB = thr_mma.partition_B(gB_nkl)
        tCgSFA = thr_mma.partition_A(gSFA_mkl)
        tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
        tCgC = thr_mma.partition_C(gC_mnl)

        # TMA partitions
        a_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape)
        tAsA, tAgA = cpasync.tma_partition(tma_atom_a, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sA, 0, 3), cute.group_modes(tCgA, 0, 3))
        b_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape)
        tBsB, tBgB = cpasync.tma_partition(tma_atom_b, block_in_cluster_coord_vmnk[1], b_cta_layout, cute.group_modes(sB, 0, 3), cute.group_modes(tCgB, 0, 3))
        tAsSFA, tAgSFA = cpasync.tma_partition(tma_atom_sfa, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sSFA, 0, 3), cute.group_modes(tCgSFA, 0, 3))
        tAsSFA, tAgSFA = cute.filter_zeros(tAsSFA), cute.filter_zeros(tAgSFA)
        sfb_cta_layout = cute.make_layout(cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape)
        tBsSFB, tBgSFB = cpasync.tma_partition(tma_atom_sfb, block_in_cluster_coord_sfb_vmnk[1], sfb_cta_layout, cute.group_modes(sSFB, 0, 3), cute.group_modes(tCgSFB, 0, 3))
        tBsSFB, tBgSFB = cute.filter_zeros(tBsSFB), cute.filter_zeros(tBgSFB)

        # Partition for MMA
        tCrA = tiled_mma.make_fragment_A(sA)
        tCrB = tiled_mma.make_fragment_B(sB)
        acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
        tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, self.num_acc_stage))

        pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)

        # Tensormap setup
        grid_dim = cute.arch.grid_dim()
        tensormap_workspace_idx = bidz * grid_dim[1] * grid_dim[0] + bidy * grid_dim[0] + bidx
        tensormap_manager = utils.TensorMapManager(utils.TensorMapUpdateMode.SMEM, self.bytes_per_tensormap)
        tensormap_a_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 0, None)].iterator)
        tensormap_b_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 1, None)].iterator)
        tensormap_sfa_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 2, None)].iterator)
        tensormap_sfb_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 3, None)].iterator)
        tensormap_c_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 4, None)].iterator)

        # Tile scheduler setup
        tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
        group_helper = utils.GroupedGemmTileSchedulerHelper(group_count, tile_sched_params, self.cluster_tile_shape_mnk, utils.create_initial_search_state())
        base_work_tile = tile_sched.initial_work_tile_info()
        group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
        initial_work_tile_info = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)

        # TMA WARP
        if warp_idx == self.tma_warp_id and initial_work_tile_info.is_valid_tile:
            work_tile = initial_work_tile_info
            tensormap_init_done = cutlass.Boolean(False)
            last_group_idx = cutlass.Int32(-1)
            ab_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_ab_stage)
            while work_tile.is_valid_tile:
                grouped_gemm_cta_tile_info = work_tile.group_search_result
                cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
                cur_group_idx = grouped_gemm_cta_tile_info.group_idx
                is_k_tile_cnt_zero = cur_k_tile_cnt == 0
                if not is_k_tile_cnt_zero:
                    is_group_changed = cur_group_idx != last_group_idx
                    if is_group_changed:
                        ps = (grouped_gemm_cta_tile_info.problem_shape_m, grouped_gemm_cta_tile_info.problem_shape_n, grouped_gemm_cta_tile_info.problem_shape_k)
                        real_tensor_a = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.a_dtype, ps, strides_abc, ptrs_abc, 0)
                        real_tensor_b = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.b_dtype, ps, strides_abc, ptrs_abc, 1)
                        real_tensor_sfa = self.make_tensor_sfasfb_for_tensormap_update(cur_group_idx, self.sf_dtype, ps, ptrs_sfasfb, 0)
                        real_tensor_sfb = self.make_tensor_sfasfb_for_tensormap_update(cur_group_idx, self.sf_dtype, ps, ptrs_sfasfb, 1)
                        if not tensormap_init_done:
                            self.tensormap_ab_init_barrier.arrive_and_wait()
                            tensormap_init_done = True
                        tensormap_manager.update_tensormap(
                            (real_tensor_a, real_tensor_b, real_tensor_sfa, real_tensor_sfb),
                            (tma_atom_a, tma_atom_b, tma_atom_sfa, tma_atom_sfb),
                            (tensormap_a_gmem_ptr, tensormap_b_gmem_ptr, tensormap_sfa_gmem_ptr, tensormap_sfb_gmem_ptr),
                            self.tma_warp_id, (tensormap_a_smem_ptr, tensormap_b_smem_ptr, tensormap_sfa_smem_ptr, tensormap_sfb_smem_ptr),
                        )
                    mma_tile_coord_mnl = (grouped_gemm_cta_tile_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape), grouped_gemm_cta_tile_info.cta_tile_idx_n, 0)
                    tAgA_slice = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
                    tBgB_slice = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
                    tAgSFA_slice = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
                    tBgSFB_slice = tBgSFB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
                    ab_producer_state.reset_count()
                    peek_ab_empty_status = cutlass.Boolean(1)
                    if ab_producer_state.count < cur_k_tile_cnt:
                        peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
                    if is_group_changed:
                        tensormap_manager.fence_tensormap_update(tensormap_a_gmem_ptr)
                        tensormap_manager.fence_tensormap_update(tensormap_b_gmem_ptr)
                        tensormap_manager.fence_tensormap_update(tensormap_sfa_gmem_ptr)
                        tensormap_manager.fence_tensormap_update(tensormap_sfb_gmem_ptr)
                    for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
                        ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
                        bar_ptr = ab_pipeline.producer_get_barrier(ab_producer_state)
                        cute.copy(tma_atom_a, tAgA_slice[(None, ab_producer_state.count)], tAsA[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=a_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_a_gmem_ptr, cute.AddressSpace.generic))
                        cute.copy(tma_atom_b, tBgB_slice[(None, ab_producer_state.count)], tBsB[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=b_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_b_gmem_ptr, cute.AddressSpace.generic))
                        cute.copy(tma_atom_sfa, tAgSFA_slice[(None, ab_producer_state.count)], tAsSFA[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=sfa_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfa_gmem_ptr, cute.AddressSpace.generic))
                        cute.copy(tma_atom_sfb, tBgSFB_slice[(None, ab_producer_state.count)], tBsSFB[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=sfb_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfb_gmem_ptr, cute.AddressSpace.generic))
                        ab_producer_state.advance()
                        peek_ab_empty_status = cutlass.Boolean(1)
                        if ab_producer_state.count < cur_k_tile_cnt:
                            peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
                else:
                    if not tensormap_init_done:
                        self.tensormap_ab_init_barrier.arrive_and_wait()
                        tensormap_init_done = True
                tile_sched.advance_to_next_work()
                base_work_tile = tile_sched.get_current_work()
                group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
                work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
                last_group_idx = cur_group_idx
            ab_pipeline.producer_tail(ab_producer_state)

        # MMA WARP
        if warp_idx == self.mma_warp_id and initial_work_tile_info.is_valid_tile:
            tensormap_manager.init_tensormap_from_atom(tma_atom_a, tensormap_a_smem_ptr, self.mma_warp_id)
            tensormap_manager.init_tensormap_from_atom(tma_atom_b, tensormap_b_smem_ptr, self.mma_warp_id)
            tensormap_manager.init_tensormap_from_atom(tma_atom_sfa, tensormap_sfa_smem_ptr, self.mma_warp_id)
            tensormap_manager.init_tensormap_from_atom(tma_atom_sfb, tensormap_sfb_smem_ptr, self.mma_warp_id)
            self.tensormap_ab_init_barrier.arrive_and_wait()
            self.tmem_alloc_barrier.arrive_and_wait()
            acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(self.acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
            sfa_tmem_ptr = cute.recast_ptr(acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base), dtype=self.sf_dtype)
            tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(tiled_mma, self.mma_tiler, self.sf_vec_size, cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)))
            tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
            sfb_tmem_ptr = cute.recast_ptr(acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base) + tcgen05.find_tmem_tensor_col_offset(tCtSFA), dtype=self.sf_dtype)
            tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(tiled_mma, self.mma_tiler, self.sf_vec_size, cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)))
            tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
            tiled_copy_s2t_sfa, tCsSFA_compact_s2t, tCtSFA_compact_s2t = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
            tiled_copy_s2t_sfb, tCsSFB_compact_s2t, tCtSFB_compact_s2t = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
            work_tile = initial_work_tile_info
            ab_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_ab_stage)
            acc_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_acc_stage)
            while work_tile.is_valid_tile:
                problem_shape_k = work_tile.group_search_result.problem_shape_k
                cur_k_tile_cnt = (problem_shape_k + self.cluster_tile_shape_mnk[2] - 1) // self.cluster_tile_shape_mnk[2]
                is_k_tile_cnt_zero = cur_k_tile_cnt == 0
                tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
                ab_consumer_state.reset_count()
                peek_ab_full_status = cutlass.Boolean(1)
                if ab_consumer_state.count < cur_k_tile_cnt and is_leader_cta:
                    peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
                if is_leader_cta and not is_k_tile_cnt_zero:
                    acc_pipeline.producer_acquire(acc_producer_state)
                tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
                for k_tile in range(cur_k_tile_cnt):
                    if is_leader_cta:
                        ab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status)
                        s2t_stage_coord = (None, None, None, None, ab_consumer_state.index)
                        cute.copy(tiled_copy_s2t_sfa, tCsSFA_compact_s2t[s2t_stage_coord], tCtSFA_compact_s2t)
                        cute.copy(tiled_copy_s2t_sfb, tCsSFB_compact_s2t[s2t_stage_coord], tCtSFB_compact_s2t)
                        num_kblocks = cute.size(tCrA, mode=[2])
                        for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                            kblock_coord = (None, None, kblock_idx, ab_consumer_state.index)
                            sf_kblock_coord = (None, None, kblock_idx)
                            tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator)
                            tiled_mma.set(tcgen05.Field.SFB, tCtSFB[sf_kblock_coord].iterator)
                            cute.gemm(tiled_mma, tCtAcc, tCrA[kblock_coord], tCrB[kblock_coord], tCtAcc)
                            tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
                        ab_pipeline.consumer_release(ab_consumer_state)
                    ab_consumer_state.advance()
                    peek_ab_full_status = cutlass.Boolean(1)
                    if ab_consumer_state.count < cur_k_tile_cnt:
                        if is_leader_cta:
                            peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
                if not is_k_tile_cnt_zero:
                    if is_leader_cta:
                        acc_pipeline.producer_commit(acc_producer_state)
                    acc_producer_state.advance()
                tile_sched.advance_to_next_work()
                base_work_tile = tile_sched.get_current_work()
                group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
                work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
            acc_pipeline.producer_tail(acc_producer_state)

        # EPILOGUE WARPS
        if warp_idx < self.mma_warp_id and initial_work_tile_info.is_valid_tile:
            tensormap_manager.init_tensormap_from_atom(tma_atom_c, tensormap_c_smem_ptr, self.epilog_warp_id[0])
            if warp_idx == self.epilog_warp_id[0]:
                cute.arch.alloc_tmem(self.num_tmem_alloc_cols, tmem_holding_buf, is_two_cta=use_2cta_instrs)
            self.tmem_alloc_barrier.arrive_and_wait()
            acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(self.acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
            epi_tidx = tidx
            tiled_copy_t2r, tTR_tAcc_base, tTR_rAcc = self.epilog_tmem_copy_and_partition(epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs)
            tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
            tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(tiled_copy_t2r, tTR_rC, epi_tidx, sC)
            tma_atom_c, bSG_sC, bSG_gC_partitioned = self.epilog_gmem_copy_and_partition(epi_tidx, tma_atom_c, tCgC, epi_tile, sC)
            work_tile = initial_work_tile_info
            acc_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_acc_stage)
            c_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, 32 * len(self.epilog_warp_id))
            c_pipeline = pipeline.PipelineTmaStore.create(num_stages=self.num_c_stage, producer_group=c_producer_group)
            last_group_idx = cutlass.Int32(-1)
            while work_tile.is_valid_tile:
                grouped_gemm_cta_tile_info = work_tile.group_search_result
                cur_group_idx = grouped_gemm_cta_tile_info.group_idx
                cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
                is_k_tile_cnt_zero = cur_k_tile_cnt == 0
                is_group_changed = cur_group_idx != last_group_idx
                if is_group_changed:
                    ps = (grouped_gemm_cta_tile_info.problem_shape_m, grouped_gemm_cta_tile_info.problem_shape_n, grouped_gemm_cta_tile_info.problem_shape_k)
                    real_tensor_c = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.c_dtype, ps, strides_abc, ptrs_abc, 2)
                    tensormap_manager.update_tensormap(((real_tensor_c),), ((tma_atom_c),), ((tensormap_c_gmem_ptr),), self.epilog_warp_id[0], (tensormap_c_smem_ptr,))
                mma_tile_coord_mnl = (grouped_gemm_cta_tile_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape), grouped_gemm_cta_tile_info.cta_tile_idx_n, 0)
                bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
                tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_consumer_state.index)]
                if not is_k_tile_cnt_zero:
                    acc_pipeline.consumer_wait(acc_consumer_state)
                tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
                bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
                if is_group_changed:
                    if warp_idx == self.epilog_warp_id[0]:
                        tensormap_manager.fence_tensormap_update(tensormap_c_gmem_ptr)
                subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
                num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
                for subtile_idx in range(subtile_cnt):
                    if not is_k_tile_cnt_zero:
                        cute.copy(tiled_copy_t2r, tTR_tAcc[(None, None, None, subtile_idx)], tTR_rAcc)
                        acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
                        tRS_rC.store(acc_vec.to(self.c_dtype))
                    else:
                        if cutlass.const_expr(self.is_nvfp4_output):
                            zeros_i8 = cute.make_rmem_tensor(cute.recast_layout(cutlass.Int8.width, self.c_dtype.width, tRS_rC.layout), cutlass.Int8)
                            zeros_i8.fill(0)
                            tRS_rC.store(cute.recast_tensor(zeros_i8, self.c_dtype).load())
                        else:
                            tRS_rC.fill(0)
                    c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
                    cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
                    cute.arch.fence_proxy(cute.arch.ProxyKind.async_shared, space=cute.arch.SharedSpace.shared_cta)
                    self.epilog_sync_barrier.arrive_and_wait()
                    if warp_idx == self.epilog_warp_id[0]:
                        cute.copy(tma_atom_c, bSG_sC[(None, c_buffer)], bSG_gC[(None, subtile_idx)], tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_c_gmem_ptr, cute.AddressSpace.generic))
                        c_pipeline.producer_commit()
                        c_pipeline.producer_acquire()
                    self.epilog_sync_barrier.arrive_and_wait()
                if not is_k_tile_cnt_zero:
                    with cute.arch.elect_one():
                        acc_pipeline.consumer_release(acc_consumer_state)
                    acc_consumer_state.advance()
                tile_sched.advance_to_next_work()
                base_work_tile = tile_sched.get_current_work()
                group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
                work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
                last_group_idx = cur_group_idx
            if warp_idx == self.epilog_warp_id[0]:
                cute.arch.relinquish_tmem_alloc_permit(is_two_cta=use_2cta_instrs)
            self.epilog_sync_barrier.arrive_and_wait()
            if warp_idx == self.epilog_warp_id[0]:
                if use_2cta_instrs:
                    cute.arch.mbarrier_arrive(tmem_dealloc_mbar_ptr, cta_rank_in_cluster ^ 1)
                    cute.arch.mbarrier_wait(tmem_dealloc_mbar_ptr, 0)
                cute.arch.dealloc_tmem(acc_tmem_ptr, self.num_tmem_alloc_cols, is_two_cta=use_2cta_instrs)
            c_pipeline.producer_tail()

    @cute.jit
    def make_tensor_abc_for_tensormap_update(self, group_idx: cutlass.Int32, dtype: Type[cutlass.Numeric],
                                              problem_shape_mnk: tuple, strides_abc: cute.Tensor,
                                              tensor_address_abc: cute.Tensor, tensor_index: int):
        """Create tensor for A/B/C tensormap update."""
        ptr_i64 = tensor_address_abc[(group_idx, tensor_index)]
        tensor_gmem_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
        strides_tensor_gmem = strides_abc[(group_idx, tensor_index, None)]
        strides_tensor_reg = cute.make_rmem_tensor(cute.make_layout(2), strides_abc.element_type)
        cute.autovec_copy(strides_tensor_gmem, strides_tensor_reg)
        stride_mn, stride_k = strides_tensor_reg[0], strides_tensor_reg[1]
        c1, c0 = cutlass.Int32(1), cutlass.Int32(0)
        if cutlass.const_expr(tensor_index == 0):
            return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[0], problem_shape_mnk[2], c1), stride=(stride_mn, stride_k, c0)))
        elif cutlass.const_expr(tensor_index == 1):
            return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[1], problem_shape_mnk[2], c1), stride=(stride_mn, stride_k, c0)))
        else:
            return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[0], problem_shape_mnk[1], c1), stride=(stride_mn, stride_k, c0)))

    @cute.jit
    def make_tensor_sfasfb_for_tensormap_update(self, group_idx: cutlass.Int32, dtype: Type[cutlass.Numeric],
                                                 problem_shape_mnk: tuple, tensor_address_sfasfb: cute.Tensor, tensor_index: int):
        """Create tensor for SFA/SFB tensormap update."""
        ptr_i64 = tensor_address_sfasfb[(group_idx, tensor_index)]
        tensor_gmem_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
        c1 = cutlass.Int32(1)
        if cutlass.const_expr(tensor_index == 0):
            sf_layout = blockscaled_utils.tile_atom_to_shape_SF((problem_shape_mnk[0], problem_shape_mnk[2], c1), self.sf_vec_size)
        else:
            sf_layout = blockscaled_utils.tile_atom_to_shape_SF((problem_shape_mnk[1], problem_shape_mnk[2], c1), self.sf_vec_size)
        return cute.make_tensor(tensor_gmem_ptr, sf_layout)

    def mainloop_s2t_copy_and_partition(self, sSF: cute.Tensor, tSF: cute.Tensor) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        """Create S2T copy for scale factors."""
        tCsSF_compact = cute.filter_zeros(sSF)
        tCtSF_compact = cute.filter_zeros(tSF)
        copy_atom_s2t = cute.make_copy_atom(tcgen05.Cp4x32x128bOp(self.cta_group), self.sf_dtype)
        tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
        thr_copy_s2t = tiled_copy_s2t.get_slice(0)
        tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
        tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t, tCsSF_compact_s2t_)
        tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
        return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t

    def epilog_tmem_copy_and_partition(self, tidx: cutlass.Int32, tAcc: cute.Tensor, gC_mnl: cute.Tensor,
                                        epi_tile: cute.Tile, use_2cta_instrs: Union[cutlass.Boolean, bool]):
        """Create T2R copy for epilogue."""
        copy_atom_t2r = sm100_utils.get_tmem_load_op(self.cta_tile_shape_mnk, self.c_layout, self.c_dtype, self.acc_dtype, epi_tile, use_2cta_instrs)
        tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
        tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)])
        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
        tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
        gC_mnl_epi = cute.flat_divide(gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile)
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        tTR_rAcc = cute.make_rmem_tensor(tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype)
        return tiled_copy_t2r, tTR_tAcc, tTR_rAcc

    def epilog_smem_copy_and_partition(self, tiled_copy_t2r: cute.TiledCopy, tTR_rC: cute.Tensor, tidx: cutlass.Int32, sC: cute.Tensor):
        """Create R2S copy for epilogue."""
        copy_atom_r2s = sm100_utils.get_smem_store_op(self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r)
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        tRS_rC = tiled_copy_r2s.retile(tTR_rC)
        return tiled_copy_r2s, tRS_rC, tRS_sC

    def epilog_gmem_copy_and_partition(self, tidx: cutlass.Int32, atom: Union[cute.CopyAtom, cute.TiledCopy],
                                        gC_mnl: cute.Tensor, epi_tile: cute.Tile, sC: cute.Tensor):
        """Create TMA store partition for epilogue."""
        gC_epi = cute.flat_divide(gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile)
        sC_for_tma = cute.group_modes(sC, 0, 2)
        gC_for_tma = cute.group_modes(gC_epi, 0, 2)
        bSG_sC, bSG_gC = cpasync.tma_partition(atom, 0, cute.make_layout(1), sC_for_tma, gC_for_tma)
        return atom, bSG_sC, bSG_gC

    @staticmethod
    def _compute_stages(tiled_mma, mma_tiler_mnk, a_dtype, b_dtype, epi_tile, c_dtype, c_layout, sf_dtype, sf_vec_size, smem_capacity, occupancy):
        """Compute pipeline stages."""
        num_acc_stage = 1  # Single accumulator stage
        num_c_stage = 2  # Default C stages
        a_smem_layout_1 = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, a_dtype, 1)
        b_smem_layout_1 = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, b_dtype, 1)
        sfa_smem_layout_1 = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
        sfb_smem_layout_1 = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
        c_smem_layout_1 = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)
        ab_bytes = (cute.size_in_bytes(a_dtype, a_smem_layout_1) + cute.size_in_bytes(b_dtype, b_smem_layout_1) +
                    cute.size_in_bytes(sf_dtype, sfa_smem_layout_1) + cute.size_in_bytes(sf_dtype, sfb_smem_layout_1))
        mbar_bytes = 1024
        c_bytes = cute.size_in_bytes(c_dtype, c_smem_layout_1) * num_c_stage
        num_ab_stage = (smem_capacity // occupancy - (mbar_bytes + c_bytes)) // ab_bytes
        remaining = smem_capacity - occupancy * ab_bytes * num_ab_stage - occupancy * (mbar_bytes + c_bytes)
        if remaining > 0:
            num_c_stage += remaining // (occupancy * cute.size_in_bytes(c_dtype, c_smem_layout_1))
        return num_acc_stage, num_ab_stage, num_c_stage

    @staticmethod
    def _compute_grid(total_num_clusters, cluster_shape_mn, max_active_clusters):
        """Compute grid and scheduler params."""
        problem_shape = (cluster_shape_mn[0], cluster_shape_mn[1], cutlass.Int32(total_num_clusters))
        tile_sched_params = utils.PersistentTileSchedulerParams(problem_shape, (*cluster_shape_mn, 1))
        grid = utils.StaticPersistentTileScheduler.get_grid_shape(tile_sched_params, max_active_clusters)
        return tile_sched_params, grid

    @staticmethod
    def _get_mbar_smem_bytes(**kwargs_stages):
        return sum(2 * 8 * stage for stage in kwargs_stages.values())

    @staticmethod
    def is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, c_dtype):
        if ab_dtype not in {cutlass.Float4E2M1FN, cutlass.Float8E5M2, cutlass.Float8E4M3FN}: return False
        if sf_vec_size not in {16, 32}: return False
        if sf_dtype not in {cutlass.Float8E8M0FNU, cutlass.Float8E4M3FN}: return False
        if sf_dtype == cutlass.Float8E4M3FN and sf_vec_size == 32: return False
        if ab_dtype in {cutlass.Float8E5M2, cutlass.Float8E4M3FN} and sf_vec_size == 16: return False
        if c_dtype not in {cutlass.Float32, cutlass.Float16, cutlass.BFloat16, cutlass.Float8E5M2, cutlass.Float8E4M3FN}: return False
        return True

    @staticmethod
    def is_valid_layouts(ab_dtype, c_dtype, a_major, b_major, c_major):
        if ab_dtype is cutlass.Float4E2M1FN and not (a_major == "k" and b_major == "k"): return False
        return True

    @staticmethod
    def is_valid_mma_tiler_and_cluster_shape(mma_tiler_mn, cluster_shape_mn):
        if mma_tiler_mn[0] not in [128, 256] or mma_tiler_mn[1] not in [128, 256]: return False
        if cluster_shape_mn[0] % (2 if mma_tiler_mn[0] == 256 else 1) != 0: return False
        is_pow2 = lambda x: x > 0 and (x & (x - 1)) == 0
        if (cluster_shape_mn[0] * cluster_shape_mn[1] > 16 or cluster_shape_mn[0] <= 0 or cluster_shape_mn[1] <= 0 or
            cluster_shape_mn[0] > 4 or cluster_shape_mn[1] > 4 or not is_pow2(cluster_shape_mn[0]) or not is_pow2(cluster_shape_mn[1])):
            return False
        return True

    @staticmethod
    def is_valid_tensor_alignment(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major):
        def check_align(dtype, is_mode0_major, shape):
            major_idx = 0 if is_mode0_major else 1
            return shape[major_idx] % (16 * 8 // dtype.width) == 0
        for m, n, k, l in problem_sizes_mnkl:
            if not check_align(ab_dtype, a_major == "m", (m, k, l)): return False
            if not check_align(ab_dtype, b_major == "n", (n, k, l)): return False
            if not check_align(c_dtype, c_major == "m", (m, n, l)): return False
        return True

    @staticmethod
    def can_implement(ab_dtype, sf_dtype, sf_vec_size, c_dtype, mma_tiler_mn, cluster_shape_mn, problem_sizes_mnkl, a_major, b_major, c_major):
        K = Sm100GroupedBlockScaledGemmKernel
        return (K.is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, c_dtype) and
                K.is_valid_layouts(ab_dtype, c_dtype, a_major, b_major, c_major) and
                K.is_valid_mma_tiler_and_cluster_shape(mma_tiler_mn, cluster_shape_mn) and
                K.is_valid_tensor_alignment(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major))


# HOST FUNCTIONS
def create_tensor_and_stride(l, mode0, mode1, is_mode0_major, dtype, is_dynamic_layout=True):
    torch_cpu = cutlass_torch.matrix(l, mode0, mode1, is_mode0_major, cutlass.Float32)
    cute_tensor, torch_tensor = cutlass_torch.cute_tensor_like(torch_cpu, dtype, is_dynamic_layout, assumed_align=16)
    stride = (1, mode0) if is_mode0_major else (mode1, 1)
    return torch_tensor.data_ptr(), torch_tensor, cute_tensor, torch_cpu, stride

def create_tensors_abc_for_all_groups(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major):
    ref_tensors, torch_tensors, cute_tensors, strides, ptrs = [], [], [], [], []
    for m, n, k, l in problem_sizes_mnkl:
        ptr_a, t_a, c_a, r_a, s_a = create_tensor_and_stride(l, m, k, a_major == "m", ab_dtype)
        ptr_b, t_b, c_b, r_b, s_b = create_tensor_and_stride(l, n, k, b_major == "n", ab_dtype)
        ptr_c, t_c, c_c, r_c, s_c = create_tensor_and_stride(l, m, n, c_major == "m", c_dtype)
        ref_tensors.append([r_a, r_b, r_c])
        ptrs.append([ptr_a, ptr_b, ptr_c])
        torch_tensors.append([t_a, t_b, t_c])
        strides.append([s_a, s_b, s_c])
        cute_tensors.append((c_a, c_b, c_c))
    return ptrs, torch_tensors, cute_tensors, strides, ref_tensors

@cute.jit
def cvt_sf_MKL_to_M32x4xrm_K4xrk_L(sf_ref_tensor: cute.Tensor, sf_mma_tensor: cute.Tensor):
    sf_mma_tensor = cute.group_modes(cute.group_modes(sf_mma_tensor, 0, 3), 1, 3)
    for i in cutlass.range(cute.size(sf_ref_tensor)):
        mkl_coord = sf_ref_tensor.layout.get_hier_coord(i)
        sf_mma_tensor[mkl_coord] = sf_ref_tensor[mkl_coord]

def create_scale_factor_tensor(l, mn, k, sf_vec_size, dtype):
    def ceil_div(a, b): return (a + b - 1) // b
    sf_k = max(1, ceil_div(k, sf_vec_size))
    ref_shape, mma_shape = (l, mn, sf_k), (l, ceil_div(mn, 128), ceil_div(sf_k, 4), 32, 4, 4)
    ref_f32 = cutlass_torch.create_and_permute_torch_tensor(ref_shape, torch.float32, permute_order=(1,2,0),
        init_type=cutlass_torch.TensorInitType.RANDOM, init_config=cutlass_torch.RandomInitConfig(min_val=1, max_val=3))
    cute_f32 = cutlass_torch.create_and_permute_torch_tensor(mma_shape, torch.float32, permute_order=(3,4,1,5,2,0),
        init_type=cutlass_torch.TensorInitType.RANDOM, init_config=cutlass_torch.RandomInitConfig(min_val=0, max_val=1))
    cvt_sf_MKL_to_M32x4xrm_K4xrk_L(from_dlpack(ref_f32), from_dlpack(cute_f32))
    cute_f32_gpu = cute_f32.cuda()
    ref_f32 = ref_f32.permute(2,0,1).unsqueeze(-1).expand(l,mn,sf_k,sf_vec_size).reshape(l,mn,sf_k*sf_vec_size).permute(1,2,0)[:,:k,:]
    cute_tensor, cute_torch = cutlass_torch.cute_tensor_like(cute_f32, dtype, is_dynamic_layout=True, assumed_align=16)
    cute_tensor = cutlass_torch.convert_cute_tensor(cute_f32_gpu, cute_tensor, dtype, is_dynamic_layout=True)
    return ref_f32, cute_torch.data_ptr(), cute_tensor, cute_torch

def create_tensors_sfasfb_for_all_groups(problem_sizes_mnkl, sf_dtype, sf_vec_size):
    ptrs, torch_tensors, cute_tensors, refs = [], [], [], []
    for m, n, k, l in problem_sizes_mnkl:
        sfa_ref, ptr_sfa, sfa_t, sfa_torch = create_scale_factor_tensor(l, m, k, sf_vec_size, sf_dtype)
        sfb_ref, ptr_sfb, sfb_t, sfb_torch = create_scale_factor_tensor(l, n, k, sf_vec_size, sf_dtype)
        ptrs.append([ptr_sfa, ptr_sfb])
        torch_tensors.append([sfa_torch, sfb_torch])
        cute_tensors.append((sfa_t, sfb_t))
        refs.append([sfa_ref, sfb_ref])
    return ptrs, torch_tensors, cute_tensors, refs

def run(num_groups, problem_sizes_mnkl, host_problem_shape_available, ab_dtype, sf_dtype, sf_vec_size,
        c_dtype, a_major, b_major, c_major, mma_tiler_mn, cluster_shape_mn, tolerance=1e-01,
        warmup_iterations=0, iterations=1, skip_ref_check=False, **kwargs):
    """Run grouped blockscaled GEMM."""
    print(f"Running Blackwell Grouped GEMM: {num_groups} groups, MMA {mma_tiler_mn}, Cluster {cluster_shape_mn}")
    if not Sm100GroupedBlockScaledGemmKernel.can_implement(ab_dtype, sf_dtype, sf_vec_size, c_dtype, mma_tiler_mn, cluster_shape_mn, problem_sizes_mnkl, a_major, b_major, c_major):
        raise TypeError("Unsupported configuration")
    if not torch.cuda.is_available():
        raise RuntimeError("GPU required")
    torch.manual_seed(2025)

    ptrs_abc, torch_abc, cute_abc, strides_abc, ref_abc = create_tensors_abc_for_all_groups(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major)
    ptrs_sfasfb, torch_sfasfb, cute_sfasfb, refs_sfasfb = create_tensors_sfasfb_for_all_groups(problem_sizes_mnkl, sf_dtype, sf_vec_size)

    # Initial tensors
    alignment = 16
    div_ab = 32 if ab_dtype == cutlass.Float4E2M1FN else 16
    div_c = 32 if c_dtype == cutlass.Float4E2M1FN else 16
    div_sf = 32 if sf_dtype == cutlass.Float4E2M1FN else 16
    min_ab = alignment * 8 // ab_dtype.width * ((div_ab + alignment * 8 // ab_dtype.width - 1) // (alignment * 8 // ab_dtype.width))
    min_c = alignment * 8 // c_dtype.width * ((div_c + alignment * 8 // c_dtype.width - 1) // (alignment * 8 // c_dtype.width))
    min_sf = alignment * 8 // sf_dtype.width * ((div_sf + alignment * 8 // sf_dtype.width - 1) // (alignment * 8 // sf_dtype.width))
    init_abc = [create_tensor_and_stride(1, min_ab, min_ab, a_major=="m", ab_dtype)[2],
                create_tensor_and_stride(1, min_ab, min_ab, b_major=="n", ab_dtype)[2],
                create_tensor_and_stride(1, min_c, min_c, c_major=="m", c_dtype)[2]]
    init_sf = [create_tensor_and_stride(1, min_sf, min_sf, a_major=="m", sf_dtype)[2],
               create_tensor_and_stride(1, min_sf, min_sf, b_major=="n", sf_dtype)[2]]

    hw = cutlass.utils.HardwareInfo()
    sm_count = hw.get_max_active_clusters(1)
    max_active = hw.get_max_active_clusters(cluster_shape_mn[0] * cluster_shape_mn[1])
    tm_shape = (sm_count, Sm100GroupedBlockScaledGemmKernel.num_tensormaps, Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8)
    tensor_of_tensormap, _ = cutlass_torch.cute_tensor_like(torch.empty(tm_shape, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False)

    gemm = Sm100GroupedBlockScaledGemmKernel(sf_vec_size, mma_tiler_mn, cluster_shape_mn)
    td, _ = cutlass_torch.cute_tensor_like(torch.tensor(problem_sizes_mnkl, dtype=torch.int32), cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
    ts, _ = cutlass_torch.cute_tensor_like(torch.tensor([[[k,1],[k,1],[n,1]] for m,n,k,l in problem_sizes_mnkl], dtype=torch.int32), cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
    ta, _ = cutlass_torch.cute_tensor_like(torch.tensor(ptrs_abc, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
    tf, _ = cutlass_torch.cute_tensor_like(torch.tensor(ptrs_sfasfb, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False, assumed_align=16)

    cta_m = 128 * cluster_shape_mn[0]
    cta_n = mma_tiler_mn[1] * cluster_shape_mn[1]
    tc = sum(((m+cta_m-1)//cta_m)*((n+cta_n-1)//cta_n) for m,n,_,_ in problem_sizes_mnkl)
    if not host_problem_shape_available: tc = max_active

    compiled = cute.compile(gemm, init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], num_groups, td, ts, ta, tf, tc, tensor_of_tensormap, max_active, options="--opt-level 2")

    if not skip_ref_check:
        compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
        print("Verifying...")
        for i, ((a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (a_t, b_t, c_t), (m, n, k, l)) in enumerate(zip(ref_abc, refs_sfasfb, cute_abc, problem_sizes_mnkl)):
            ref = torch.einsum("mkl,nkl->mnl", torch.einsum("mkl,mkl->mkl", a_ref, sfa_ref), torch.einsum("nkl,nkl->nkl", b_ref, sfb_ref))
            c_gpu = c_ref.cuda()
            cute.testing.convert(c_t, from_dlpack(c_gpu, assumed_align=16).mark_layout_dynamic(leading_dim=(1 if c_major=="n" else 0)))
            c_cpu = c_gpu.cpu()
            if c_dtype in (cutlass.Float32, cutlass.Float16, cutlass.BFloat16):
                torch.testing.assert_close(c_cpu, ref, atol=tolerance, rtol=1e-02)
            elif c_dtype in (cutlass.Float8E5M2, cutlass.Float8E4M3FN):
                ref_f8_ = torch.empty(*(l,m,n), dtype=torch.uint8, device="cuda").permute(1,2,0)
                ref_f8 = from_dlpack(ref_f8_, assumed_align=16).mark_layout_dynamic(leading_dim=1)
                ref_f8.element_type = c_dtype
                ref_dev = ref.permute(2,0,1).contiguous().permute(1,2,0).cuda()
                ref_t = from_dlpack(ref_dev, assumed_align=16).mark_layout_dynamic(leading_dim=1)
                cute.testing.convert(ref_t, ref_f8)
                cute.testing.convert(ref_f8, ref_t)
                torch.testing.assert_close(c_cpu, ref_dev.cpu(), atol=tolerance, rtol=1e-02)

    for _ in range(warmup_iterations):
        compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
    torch.cuda.synchronize()
    start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(iterations):
        compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
    end.record()
    torch.cuda.synchronize()
    exec_time = start.elapsed_time(end) * 1000 / iterations
    fmas = sum(m*n*k for m,n,k,_ in problem_sizes_mnkl)
    gflops = 2 * fmas / 1e9 / (exec_time / 1e6)
    print(f"Runtime: {exec_time/1000:.3f}ms, GFLOPS: {gflops:.1f}")
    return exec_time


if __name__ == "__main__":
    def parse_ints(s): return tuple(int(x.strip()) for x in s.split(","))
    def parse_tuples(s):
        tuples = s.strip("()").split("),(")
        return [tuple(int(x.strip()) for x in t.split(",")) for t in tuples]

    parser = argparse.ArgumentParser(description="Grouped GEMM on Blackwell")
    parser.add_argument("--num_groups", type=int, default=2)
    parser.add_argument("--problem_sizes_mnkl", type=parse_tuples, default=((128,128,128,1),(128,128,128,1)))
    parser.add_argument("--mma_tiler_mn", type=parse_ints, default=(128,128))
    parser.add_argument("--host_problem_shape_available", action="store_true")
    parser.add_argument("--cluster_shape_mn", type=parse_ints, default=(1,1))
    parser.add_argument("--ab_dtype", type=cutlass.dtype, default=cutlass.Float4E2M1FN)
    parser.add_argument("--sf_dtype", type=cutlass.dtype, default=cutlass.Float8E8M0FNU)
    parser.add_argument("--sf_vec_size", type=int, default=16)
    parser.add_argument("--c_dtype", type=cutlass.dtype, default=cutlass.Float16)
    parser.add_argument("--a_major", choices=["k","m"], default="k")
    parser.add_argument("--b_major", choices=["k","n"], default="k")
    parser.add_argument("--c_major", choices=["n","m"], default="n")
    parser.add_argument("--tolerance", type=float, default=1e-01)
    parser.add_argument("--warmup_iterations", type=int, default=0)
    parser.add_argument("--iterations", type=int, default=1)
    parser.add_argument("--skip_ref_check", action="store_true")
    args = parser.parse_args()
    if len(args.problem_sizes_mnkl) != args.num_groups:
        parser.error("problem_sizes_mnkl must contain num_groups tuples")
    for _, _, _, l in args.problem_sizes_mnkl:
        if l != 1: parser.error("l must be 1")
    run(args.num_groups, args.problem_sizes_mnkl, args.host_problem_shape_available, args.ab_dtype,
        args.sf_dtype, args.sf_vec_size, args.c_dtype, args.a_major, args.b_major, args.c_major,
        args.mma_tiler_mn, args.cluster_shape_mn, args.tolerance, args.warmup_iterations,
        args.iterations, args.skip_ref_check)
    print("PASS")


# AUTO-TESTING ENTRY POINT
try:
    from task import input_t, output_t
except ImportError:
    input_t = Tuple[List[Tuple[torch.Tensor, torch.Tensor, torch.Tensor]], List[Tuple[torch.Tensor, torch.Tensor]], List[Tuple[torch.Tensor, torch.Tensor]], List[Tuple[int, int, int, int]]]
    output_t = List[torch.Tensor]

_base_configs = {}
_cache, _ptr_cache, _last_call = {}, {}, None
_single_kernel_cache = {}  # Cache for single-group kernels
_graph_cache = {}  # Cache for CUDA graphs

def _sort_groups_by_k_desc(ps, at, sr):
    """Sort groups by K dimension in descending order."""
    indexed = [(i, ps[i][2]) for i in range(len(ps))]
    sorted_indexed = sorted(indexed, key=lambda x: -x[1])  # Descending K
    sorted_indices = [x[0] for x in sorted_indexed]
    inverse_indices = [0] * len(sorted_indices)
    for new_idx, old_idx in enumerate(sorted_indices):
        inverse_indices[old_idx] = new_idx
    return sorted_indices, inverse_indices

def _sort_groups_by_work(ps, at, sr):
    """Sort groups by total work (M*N*K) in descending order for better load balancing."""
    indexed = [(i, ps[i][0] * ps[i][1] * ps[i][2]) for i in range(len(ps))]
    sorted_indexed = sorted(indexed, key=lambda x: -x[1])  # Descending order
    sorted_indices = [x[0] for x in sorted_indexed]
    inverse_indices = [0] * len(sorted_indices)
    for new_idx, old_idx in enumerate(sorted_indices):
        inverse_indices[old_idx] = new_idx
    return sorted_indices, inverse_indices

def _sort_groups_by_nk(ps, at, sr):
    indexed = [(i, (ps[i][1], ps[i][2])) for i in range(len(ps))]
    sorted_indexed = sorted(indexed, key=lambda x: x[1])
    sorted_indices = [x[0] for x in sorted_indexed]
    inverse_indices = [0] * len(sorted_indices)
    for new_idx, old_idx in enumerate(sorted_indices):
        inverse_indices[old_idx] = new_idx
    return sorted_indices, inverse_indices

def _get_base(config_key):
    """Get or create base tensors for a given configuration."""
    global _base_configs
    MMA_M, MMA_N, CLUSTER_M, CLUSTER_N = config_key
    if config_key not in _base_configs:
        ab, sf, c = cutlass.Float4E2M1FN, cutlass.Float8E4M3FN, cutlass.Float16
        abc = [create_tensor_and_stride(1,64,64,False,ab)[2], create_tensor_and_stride(1,64,64,False,ab)[2], create_tensor_and_stride(1,16,16,False,c)[2]]
        sft = [create_tensor_and_stride(1,16,16,False,sf)[2], create_tensor_and_stride(1,16,16,False,sf)[2]]
        hw = cutlass.utils.HardwareInfo()
        CLUSTER_SIZE = CLUSTER_M * CLUSTER_N
        sm = hw.get_max_active_clusters(CLUSTER_SIZE)
        tm = torch.zeros((sm*CLUSTER_SIZE,5,16), dtype=torch.int64, device="cuda")
        tmap, _ = cutlass_torch.cute_tensor_like(tm, cutlass.Int64, is_dynamic_layout=False)
        _base_configs[config_key] = {"abc": abc, "sf": sft, "tmap": tmap, "max": sm, "cluster_size": CLUSTER_SIZE}
    return _base_configs[config_key]


def _get_single_group_kernel(m, n, k, config_key):
    """Get or compile a single-group kernel for the given problem size."""
    global _single_kernel_cache
    MMA_M, MMA_N, CLUSTER_M, CLUSTER_N = config_key
    CLUSTER_TILE_M = 128 * CLUSTER_M
    CLUSTER_TILE_N = MMA_N * CLUSTER_N
    
    # Cache key based on tile counts (not exact size) for reuse
    tm = (m + CLUSTER_TILE_M - 1) // CLUSTER_TILE_M
    tn = (n + CLUSTER_TILE_N - 1) // CLUSTER_TILE_N
    tk = k
    cache_key = (tm, tn, tk, config_key)
    
    if cache_key not in _single_kernel_cache:
        b = _get_base(config_key)
        ps = [(m, n, k, 1)]
        st = [[[k,1],[k,1],[n,1]]]
        tc = tm * tn
        
        d = torch.tensor(ps, dtype=torch.int32, device="cuda")
        s = torch.tensor(st, dtype=torch.int32, device="cuda")
        pa = torch.zeros((1, 3), dtype=torch.int64, device="cuda")
        pf = torch.zeros((1, 2), dtype=torch.int64, device="cuda")
        
        td, _ = cutlass_torch.cute_tensor_like(d, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
        ts, _ = cutlass_torch.cute_tensor_like(s, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
        ta, ta_torch = cutlass_torch.cute_tensor_like(pa, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
        tf, tf_torch = cutlass_torch.cute_tensor_like(pf, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
        
        g = Sm100GroupedBlockScaledGemmKernel(16, (MMA_M, MMA_N), (CLUSTER_M, CLUSTER_N), occupancy=1)
        kernel = cute.compile(g, b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1], 
                             1, td, ts, ta, tf, tc, b["tmap"], b["max"], options="--opt-level 3")
        _single_kernel_cache[cache_key] = {
            "kernel": kernel, "td": td, "ts": ts, "ta": ta, "tf": tf, 
            "ta_torch": ta_torch, "tf_torch": tf_torch, "base": b
        }
    return _single_kernel_cache[cache_key]


def _run_separate_kernels(at, sr, ps, config_key):
    """Run separate single-group kernels for each group."""
    ng = len(ps)
    outs = []
    
    for i in range(ng):
        m, n, k, l = ps[i]
        a, b_, c = at[i]
        sa, sb = sr[i]
        
        r = _get_single_group_kernel(m, n, k, config_key)
        
        # Update pointers
        r["ta_torch"][0, 0] = a.data_ptr()
        r["ta_torch"][0, 1] = b_.data_ptr()
        r["ta_torch"][0, 2] = c.data_ptr()
        r["tf_torch"][0, 0] = sa.data_ptr()
        r["tf_torch"][0, 1] = sb.data_ptr()
        
        # Run kernel for this group
        b = r["base"]
        r["kernel"](b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1],
                   r["td"], r["ts"], r["ta"], r["tf"], b["tmap"])
        outs.append(c)
    
    return outs

def custom_kernel(data: input_t) -> output_t:
    global _cache, _ptr_cache, _last_call
    at, _, sr, ps = data
    ng = len(ps)
    
    # Default config
    MMA_M, MMA_N, CLUSTER_M, CLUSTER_N, OCCUPANCY = 128, 128, 1, 1, 1
    
    CLUSTER_TILE_M, CLUSTER_TILE_N = 128 * CLUSTER_M, MMA_N * CLUSTER_N
    config_key = (MMA_M, MMA_N, CLUSTER_M, CLUSTER_N)
    
    if _last_call is not None:
        last_ng, last_at, last_sr, cached = _last_call
        if last_ng == ng and last_at is at and last_sr is sr:
            kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap = cached["args"]
            kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
            return cached["outs"]
    
    b = _get_base(config_key)
    ptr_at = tuple((a.data_ptr(), b_.data_ptr(), c.data_ptr()) for (a, b_, c) in at)
    ptr_sr = tuple((sa.data_ptr(), sb.data_ptr()) for (sa, sb) in sr)
    ptr_key = (ng, ptr_at, ptr_sr, tuple(tuple(x) for x in ps), config_key)
    
    if ptr_key in _ptr_cache:
        r, outs, cached_args = _ptr_cache[ptr_key]
        kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap = cached_args
        kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
        _last_call = (ng, at, sr, {"args": cached_args, "outs": outs})
        return outs
    
    # Keep groups in original order - let the scheduler handle distribution
    sorted_indices = list(range(ng))
    ps_sorted = [ps[i] for i in sorted_indices]
    at_sorted = [at[i] for i in sorted_indices]
    sr_sorted = [sr[i] for i in sorted_indices]
    tc = sum(((m+CLUSTER_TILE_M-1)//CLUSTER_TILE_M)*((n+CLUSTER_TILE_N-1)//CLUSTER_TILE_N) for m,n,_,_ in ps_sorted)
    key = (ng, tc, tuple(tuple(x) for x in ps_sorted), config_key)
    ap = [[a.data_ptr(), b_.data_ptr(), c.data_ptr()] for (a, b_, c) in at_sorted]
    sp = [[sa.data_ptr(), sb.data_ptr()] for (sa, sb) in sr_sorted]
    st = [[[k,1],[k,1],[n,1]] for m,n,k,l in ps_sorted]
    
    if key not in _cache:
        d = torch.tensor(ps_sorted, dtype=torch.int32, device="cuda")
        s = torch.tensor(st, dtype=torch.int32, device="cuda")
        pa = torch.tensor(ap, dtype=torch.int64, device="cuda")
        pf = torch.tensor(sp, dtype=torch.int64, device="cuda")
        td, _ = cutlass_torch.cute_tensor_like(d, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
        ts, _ = cutlass_torch.cute_tensor_like(s, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
        ta, ta_torch = cutlass_torch.cute_tensor_like(pa, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
        tf, tf_torch = cutlass_torch.cute_tensor_like(pf, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
        g = Sm100GroupedBlockScaledGemmKernel(16, (MMA_M, MMA_N), (CLUSTER_M, CLUSTER_N), occupancy=1, max_ab_stages=0)
        k = cute.compile(g, b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1], ng, td, ts, ta, tf, tc, b["tmap"], b["max"], options="--opt-level 3")
        _cache[key] = {"k": k, "td": td, "ts": ts, "ta": ta, "tf": tf, "ta_torch": ta_torch, "tf_torch": tf_torch}
    
    r = _cache[key]
    if "ap_cpu" not in r:
        r["ap_cpu"] = torch.zeros(ng*3, dtype=torch.int64, pin_memory=True)
        r["sp_cpu"] = torch.zeros(ng*2, dtype=torch.int64, pin_memory=True)
    
    # Direct copy from pinned CPU to cute tensor's underlying torch tensor
    ap_cpu, sp_cpu = r["ap_cpu"], r["sp_cpu"]
    idx = 0
    for (a_, b_, c_) in at_sorted:
        ap_cpu[idx], ap_cpu[idx+1], ap_cpu[idx+2] = a_.data_ptr(), b_.data_ptr(), c_.data_ptr()
        idx += 3
    idx = 0
    for (sa, sb) in sr_sorted:
        sp_cpu[idx], sp_cpu[idx+1] = sa.data_ptr(), sb.data_ptr()
        idx += 2
    
    # Async copy: pinned CPU -> GPU cute tensor (don't wait for completion)
    r["ta_torch"].view(-1).copy_(ap_cpu, non_blocking=True)
    r["tf_torch"].view(-1).copy_(sp_cpu, non_blocking=True)
    
    outs = [at[i][2] for i in range(ng)]
    kernel = r["k"]
    abc0, abc1, abc2 = b["abc"][0], b["abc"][1], b["abc"][2]
    sf0, sf1 = b["sf"][0], b["sf"][1]
    td, ts, ta, tf, tmap = r["td"], r["ts"], r["ta"], r["tf"], b["tmap"]
    
    # Store args for caching
    cached_args = (kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
    _ptr_cache[ptr_key] = (r, outs, cached_args)
    _last_call = (ng, at, sr, {"args": cached_args, "outs": outs})
    
    # Direct kernel call
    kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
    return outs

print("main")
print("main end")
scrolls · 1091 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 409213.

⋯ 66 unchanged lines
self.tensormap_update_mode = utils.TensorMapUpdateMode.SMEM
self.occupancy = occupancy
self.max_ab_stages = max_ab_stages # 0 = auto, >0 = limit stages
- self.epilog_warp_id = (0, 1, 2, 3)
+ self.epilog_warp_id = (0, 1, 2, 3) # 4 epilogue warps
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_cta = 32 * len((self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id))
⋯ 555 unchanged lines
def _compute_stages(tiled_mma, mma_tiler_mnk, a_dtype, b_dtype, epi_tile, c_dtype, c_layout, sf_dtype, sf_vec_size, smem_capacity, occupancy):
"""Compute pipeline stages."""
num_acc_stage = 1 # Single accumulator stage
- num_c_stage = 2
+ num_c_stage = 2 # Default C stages
a_smem_layout_1 = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, a_dtype, 1)
b_smem_layout_1 = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, b_dtype, 1)
sfa_smem_layout_1 = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
⋯ 431 unchanged lines
sp_cpu[idx], sp_cpu[idx+1] = sa.data_ptr(), sb.data_ptr()
idx += 2
- # Single copy: pinned CPU -> GPU cute tensor
+ # Async copy: pinned CPU -> GPU cute tensor (don't wait for completion)
r["ta_torch"].view(-1).copy_(ap_cpu, non_blocking=True)
r["tf_torch"].view(-1).copy_(sp_cpu, non_blocking=True)
scrolls · 27 diff lines total

Best evidence level for this revision: reported

JSON