Skip to content
KernelIndex
Search⌘K

submission 483112

John · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_manual.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-483112?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.4µs
#123 of 310
2026-02-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bba92b1be1305e252b3d4ac23dcce37265be3ef720866146250af5d923936b73
license declaredunknown
license concludedunknown
authorsJohn
imported2026-08-15

Techniques

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

fused-epilogueepi_tile = sm100_utils.compute_epilogue_tile_shape(
mbarrierself.epilog_sync_barrier = pipeline.NamedBarrier(
persistent-kerneltile_sched_params: utils.PersistentTileSchedulerParams,
shared-memorysmem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
tcgen05tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission_manual.py2574 lines
#!POPCORN leaderboard nvfp4_group_gemm

import cutlass
import cutlass.cute as cute
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
from cutlass.cute.runtime import make_ptr

from inspect import isclass

import functools
from typing import Tuple, List, Type, Union

import torch
from task import input_t, output_t

# Kernel configuration parameters
# Size of tma descriptor in bytes
bytes_per_tensormap = 128

# Number of tensormaps: a, b, sfa, sfb
num_tensormaps = 4
# Tile sizes for M, N, K dimensions
mma_tiler_mnk = (128, 128, 256)
# Shape of the K dimension for the MMA instruction
mma_inst_shape_k = 64
mma_inst_tile_k = 4

# FP4 data type for A and B
ab_dtype = cutlass.Float4E2M1FN
# FP8 data type for scale factors
sf_dtype = cutlass.Float8E4M3FN
# FP16 output type
c_dtype = cutlass.Float16
# Scale factor block size (16 elements share one scale)
sf_vec_size = 16

# Total number of columns in tmem
num_tmem_alloc_cols = 512

# Staging knobs for tuning.
# max_num_ab_stage=0 means no cap.
# force_num_c_stage=0 means keep computed value.
max_num_ab_stage = 6
force_num_c_stage = 0

# Persistent scheduler knobs.
# tile_swizzle_size=1 means no swizzle.
tile_swizzle_size = 1
# Only used when tile_swizzle_size > 1.
tile_raster_along_m = True

smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")

buffer_align_bytes = 1024

"""
 A sensible default is (1, 1)—that means no CTA clustering. Each cluster is just one CTA, so:

  - simplest scheduling,
  - no multicast,
  - lowest coordination overhead,
  - good baseline for correctness and first‑time tuning.

  You typically consider larger cluster shapes when:

  - M and N are both large (more parallelism),
  - you want multicast reuse of A/B across CTAs,
  - and you’re targeting performance on SM100.

"""
cluster_shape_mn = (1, 1)

class Sm100GroupedBlockScaledGemmKernel:
    # Size of smem we reserved for mbarrier, tensor memory management and tensormap update
    reserved_smem_bytes = 1024

    # bytes_per_tensormap is the size of one TMA descriptor (tensormap).
    # In this kernel it’s fixed at 128 bytes, which is what the hardware expects for a TMA descriptor on SM100.
    bytes_per_tensormap = 128

    # 5 TMA descriptor
    # A, B, sfA, sfB, C
    num_tensormaps = 5

    tensor_memory_management_bytes = 12

    def __init__(
            self,
            mma_tiler_mn: Tuple[int, int],
            cluster_shape_mn: Tuple[int, int],
        ):
        self.acc_dtype = cutlass.Float32
        self.use_2cta_instrs = mma_tiler_mn[0] == 256
        self.cluster_shape_mn = cluster_shape_mn
        # K dimension is deferred in _setup_attributes
        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.epilog_warp_id = (
            0,
            1,
            2,
            3,
        )
        self.mma_warp_id = 4
        self.tma_warp_id = 5

        self.occupancy = 1

        self.threads_per_cta = 32 * len(
            (self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
        )

        # bytes_per_tensormap is in bytes, so dividing by 8 converts bytes → number of int64 elements
        self.size_tensormap_in_i64 = (
            Sm100GroupedBlockScaledGemmKernel.num_tensormaps
            * Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap
            // 8
        )

        # Set barrier for epilogue sync and tmem ptr sync
        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)),
        )
        # Barrier used by MMA/TMA warps to signal A/B tensormap initialization completion
        self.tensormap_ab_init_barrier = pipeline.NamedBarrier(
            barrier_id=3,
            num_threads=64,
        )

    @cute.kernel
    def ported_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,
        ptrs_abc: cute.Tensor,
        ptrs_sfasfb: cute.Tensor,
        tensormaps: cute.Tensor,
    ):
        warp_idx = cute.arch.warp_idx()
        # Hinting to the compiler this is identical across warp, enables warp level control flow, no divergent branches, better  codegen.
        warp_idx = cute.arch.make_warp_uniform(warp_idx)

        # It runs only in the dedicated TMA warp and prefetches the TMA descriptors for A/B/SFA/SFB/C into the hardware’s descriptor cache.
        # That reduces latency when the warp later issues TMA loads/stores using those descriptors.
        # This is just warming the cache; it doesn’t move any actual tensor data yet.
        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

        # Figuring out where the current CTA and thread sit inside the cluster so it can pick the right slice of work.
        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

        # Noting, best practice to hint warp level variables.
        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
        )
        # coord inside cta
        tidx, _, _ = cute.arch.thread_idx()

        # Shared memory is just a raw byte buffer per CTA.
        # You need a structured layout to carve it into named regions (tensormap buffer, mbarriers, sA/sB/sC tiles, etc.).
        # @cute.struct lets you describe that layout in one place with sizes and alignment
        # so later code can refer to storage.sA, storage.sB, etc. without manual pointer math.
        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
            + Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
        )
        tensormap_sfa_smem_ptr = (
            tensormap_b_smem_ptr
            + Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
        )
        tensormap_sfb_smem_ptr = (
            tensormap_sfa_smem_ptr
            + Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
        )
        tensormap_c_smem_ptr = (
            tensormap_sfb_smem_ptr
            + Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
        )

        tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
        tmem_holding_buf = storage.tmem_holding_buf

        # Initialize mainloop ab_pipeline (barrier) and states
        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,
        )

        # Initialize acc_pipeline (barrier) and states
        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,
        )

        # Tensor memory dealloc barrier init
        if use_2cta_instrs:
            if warp_idx == self.tma_warp_id:
                num_tmem_dealloc_threads = 32
                with cute.arch.elect_one():
                    cute.arch.mbarrier_init(
                        tmem_dealloc_mbar_ptr, num_tmem_dealloc_threads
                    )

        # Cluster arrive after barrier init
        #  The pipeline uses mbarriers to coordinate between CTAs in a cluster.
        # At startup, those barriers are in an “empty/not arrived” state. pipeline_init_arrive(...) marks the initial arrival
        # so the other CTAs (or later code) that do a pipeline_init_wait(...) don’t deadlock on a barrier
        # that was never initialized/signaled. It’s basically “prime the barrier state” for cluster‑wide use.
        pipeline_init_arrive(cluster_shape_mn=self.cluster_shape_mn, is_relaxed=True)

        #
        # Setup smem tensor A/B/SFA/SFB/C
        # The swizzle are detemrined by a blackwell helper
        # it’s chosen to make the major dimension align to 128/64/32 bytes (1024/512/256 bits), which is optimal for TMA and SMEM access
        #
        sC = storage.sC.get_tensor(
            c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
        )
        # (MMA, MMA_M, MMA_K, STAGE)
        sA = storage.sA.get_tensor(
            a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
        )
        # (MMA, MMA_N, MMA_K, STAGE)
        sB = storage.sB.get_tensor(
            b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
        )
        # (MMA, MMA_M, MMA_K, STAGE)
        sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
        # (MMA, MMA_N, MMA_K, STAGE)
        sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)

        # Stage small per-group metadata into shared memory once per CTA to avoid
        # repeatedly reloading it from global memory in persistent loops.
        problem_sizes_smem = storage.problem_sizes_buffer.get_tensor(
            cute.make_layout((group_count, 4), stride=(4, 1))
        )
        ptrs_abc_smem = storage.ptrs_abc_buffer.get_tensor(
            cute.make_layout((group_count, 3), stride=(3, 1))
        )
        ptrs_sfasfb_smem = storage.ptrs_sfasfb_buffer.get_tensor(
            cute.make_layout((group_count, 2), stride=(2, 1))
        )
        group_tile_prefix_end = storage.group_tile_prefix_end.get_tensor(
            cute.make_layout(group_count, stride=1)
        )
        group_cluster_count_m = storage.group_cluster_count_m.get_tensor(
            cute.make_layout(group_count, stride=1)
        )
        group_cluster_count_k = storage.group_cluster_count_k.get_tensor(
            cute.make_layout(group_count, stride=1)
        )
        for linear_idx in cutlass.range(tidx, group_count * 4, self.threads_per_cta):
            group_idx = linear_idx // 4
            mode_idx = linear_idx % 4
            problem_sizes_smem[(group_idx, mode_idx)] = problem_sizes_mnkl[
                (group_idx, mode_idx)
            ]
        for linear_idx in cutlass.range(tidx, group_count * 3, self.threads_per_cta):
            group_idx = linear_idx // 3
            mode_idx = linear_idx % 3
            ptrs_abc_smem[(group_idx, mode_idx)] = ptrs_abc[(group_idx, mode_idx)]
        for linear_idx in cutlass.range(tidx, group_count * 2, self.threads_per_cta):
            group_idx = linear_idx // 2
            mode_idx = linear_idx % 2
            ptrs_sfasfb_smem[(group_idx, mode_idx)] = ptrs_sfasfb[
                (group_idx, mode_idx)
            ]
        if tidx == 0:
            tile_prefix_end = cutlass.Int32(0)
            for group_idx in cutlass.range(group_count):
                problem_m = problem_sizes_smem[(group_idx, 0)]
                problem_n = problem_sizes_smem[(group_idx, 1)]
                problem_k = problem_sizes_smem[(group_idx, 2)]
                cluster_count_m = (
                    problem_m + self.cluster_tile_shape_mnk[0] - 1
                ) // self.cluster_tile_shape_mnk[0]
                cluster_count_n = (
                    problem_n + self.cluster_tile_shape_mnk[1] - 1
                ) // self.cluster_tile_shape_mnk[1]
                cluster_count_k = (
                    problem_k + self.cluster_tile_shape_mnk[2] - 1
                ) // self.cluster_tile_shape_mnk[2]
                tile_prefix_end += cluster_count_m * cluster_count_n
                group_tile_prefix_end[group_idx] = tile_prefix_end
                group_cluster_count_m[group_idx] = cluster_count_m
                group_cluster_count_k[group_idx] = cluster_count_k
        cute.arch.barrier()


        #
        # Compute multicast mask for A/B/SFA/SFB buffer full
        #
        a_full_mcast_mask = None
        b_full_mcast_mask = None
        sfa_full_mcast_mask = None
        sfb_full_mcast_mask = 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
        # Level = (which tile of A/B/C this CTA works on
        #
        # (bM, bK, RestM, RestK, RestL)
        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        # (bN, bK, RestN, RestK, RestL)
        gB_nkl = cute.local_tile(
            mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
        )
        # (bM, bK, RestM, RestK, RestL)
        gSFA_mkl = cute.local_tile(
            mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        # (bN, bK, RestN, RestK, RestL)
        gSFB_nkl = cute.local_tile(
            mSFB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
        )
        # (bM, bN, RestM, RestN, RestL)
        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
        )


        #
        # Partition global tensor for TiledMMA_A/B/C
        # Level = what each warp/lane loads/computes
        #
        thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
        thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
        # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
        tCgA = thr_mma.partition_A(gA_mkl)
        # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
        tCgB = thr_mma.partition_B(gB_nkl)
        # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
        tCgSFA = thr_mma.partition_A(gSFA_mkl)
        # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
        tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
        # (MMA, MMA_M, MMA_N, RestM, RestN, RestL)
        tCgC = thr_mma.partition_C(gC_mnl)


        """
        Yes—TMA also needs partitioning, but for a different reason than MMA.

        Why partition for TMA:

        - TMA moves tiles from GMEM → SMEM.
        - In clustered CTAs, different CTAs are responsible for different slices of those tiles (and some are multicast).
        - cpasync.tma_partition(...) computes the per‑CTA source (tAgA/tBgB) and destination (tAsA/tBsB) slices for the TMA operation.

        So:

        - MMA partitioning is within a CTA (per‑warp/per‑thread).
        - TMA partitioning is across CTAs in a cluster, deciding which CTA loads which part and where it lands in SMEM.

        """
        #
        # Partition global/shared tensor for TMA load A/B
        #
        # TMA load A partition_S/D
        a_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestM, RestK, RestL)
        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),
        )
        # TMA load B partition_S/D
        b_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestN, RestK, RestL)
        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),
        )

        #  TMA Load SFA partition_S/D
        sfa_cta_layout = a_cta_layout
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestM, RestK, RestL)
        tAsSFA, tAgSFA = cute.nvgpu.cpasync.tma_partition(
            tma_atom_sfa,
            block_in_cluster_coord_vmnk[2],
            sfa_cta_layout,
            cute.group_modes(sSFA, 0, 3),
            cute.group_modes(tCgSFA, 0, 3),
        )
        tAsSFA = cute.filter_zeros(tAsSFA)
        tAgSFA = cute.filter_zeros(tAgSFA)

        # TMA Load SFB partition_S/D
        sfb_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestN, RestK, RestL)
        tBsSFB, tBgSFB = cute.nvgpu.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 = cute.filter_zeros(tBsSFB)
        tBgSFB = cute.filter_zeros(tBgSFB)


        #
        # Partition shared/tensor memory tensor for TiledMMA_A/B/C
        #
        # (MMA, MMA_M, MMA_K, STAGE)
        tCrA = tiled_mma.make_fragment_A(sA)
        # (MMA, MMA_N, MMA_K, STAGE)
        tCrB = tiled_mma.make_fragment_B(sB)
        # (MMA, MMA_M, MMA_N)
        acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
        # (MMA, MMA_M, MMA_N, STAGE)
        tCtAcc_fake = tiled_mma.make_fragment_C(
            cute.append(acc_shape, self.num_acc_stage)
        )

        #
        # Cluster wait before tensor memory alloc
        # pipeline_init_wait(...) is the cluster‑wide “wait for initialization” barrier.
        # It makes each CTA in the cluster wait until all CTAs have finished pipeline_init_arrive(...) (i.e., the pipeline barriers are initialized and in a consistent state).
        # This prevents any CTA from using the pipeline/barriers before others have finished setting them up, which avoids deadlocks or incorrect synchronization.
        #
        pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)


        #
        # Get tensormap buffer address
        # compute the pointer(s) to the TMA descriptor storage for this CTA’s workspace, so the kernel can read/update those descriptors.
        #
        grid_dim = cute.arch.grid_dim()
        tensormap_workspace_idx = (
            bidz * grid_dim[1] * grid_dim[0] + bidy * grid_dim[0] + bidx
        )

        # The host allocates tensormap_cute_tensor, but the kernel creates/updates the actual descriptor contents at runtime.
        # tensormap_cute_tensor is the global-memory backing store for all TMA descriptors.
        # It’s passed into the kernel so each CTA can compute its slice of that tensor and then read/update the descriptors for A/B/SF/C at runtime.
        tensormap_manager = utils.TensorMapManager(
            utils.TensorMapUpdateMode.SMEM,
            Sm100GroupedBlockScaledGemmKernel.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
        )

        #
        # Specialized TMA load warp
        #
        if warp_idx == self.tma_warp_id:
            #
            # Persistent tile scheduling loop
            #
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), grid_dim
            )
            tensormap_init_done = cutlass.Boolean(False)
            # group index of last tile
            last_group_idx = cutlass.Int32(-1)
            tma_desc_ptr_a = tensormap_manager.get_tensormap_ptr(
                tensormap_a_gmem_ptr,
                cute.AddressSpace.generic,
            )
            tma_desc_ptr_b = tensormap_manager.get_tensormap_ptr(
                tensormap_b_gmem_ptr,
                cute.AddressSpace.generic,
            )
            tma_desc_ptr_sfa = tensormap_manager.get_tensormap_ptr(
                tensormap_sfa_gmem_ptr,
                cute.AddressSpace.generic,
            )
            tma_desc_ptr_sfb = tensormap_manager.get_tensormap_ptr(
                tensormap_sfb_gmem_ptr,
                cute.AddressSpace.generic,
            )

            work_tile = tile_sched.initial_work_tile_info()
            tma_search_group_idx = cutlass.Int32(0)
            tma_search_tile_start = cutlass.Int32(0)
            tma_search_tile_end = group_tile_prefix_end[tma_search_group_idx]

            # If producer and consumer use different num_stages, they’ll index different ring sizes and the barrier protocol breaks (stages will be marked full/empty out of sync).
            # In practice, that would deadlock or corrupt data. The pipeline API assumes both sides are created with the same stage count.
            ab_producer_state = pipeline.make_pipeline_state(
                pipeline.PipelineUserType.Producer, self.num_ab_stage
            )

            while work_tile.is_valid_tile:
                cur_tile_coord = work_tile.tile_idx
                cur_linear_idx = cur_tile_coord[2]
                while (
                    cur_linear_idx >= tma_search_tile_end
                    and tma_search_group_idx + 1 < group_count
                ):
                    tma_search_group_idx += 1
                    tma_search_tile_start = tma_search_tile_end
                    tma_search_tile_end = group_tile_prefix_end[tma_search_group_idx]
                cur_group_idx = tma_search_group_idx
                cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]
                cluster_tile_idx = cur_linear_idx - tma_search_tile_start
                cluster_count_m = group_cluster_count_m[cur_group_idx]
                cta_tile_idx_m = (
                    (cluster_tile_idx % cluster_count_m) * self.cluster_shape_mn[0]
                    + cur_tile_coord[0]
                )
                cta_tile_idx_n = (
                    (cluster_tile_idx // cluster_count_m) * self.cluster_shape_mn[1]
                    + cur_tile_coord[1]
                )
                problem_shape_m = problem_sizes_smem[(cur_group_idx, 0)]
                problem_shape_n = problem_sizes_smem[(cur_group_idx, 1)]
                problem_shape_k = problem_sizes_smem[(cur_group_idx, 2)]
                is_group_changed = cur_group_idx != last_group_idx
                # skip tensormap update if we're working on the same group
                if is_group_changed:
                    real_tensor_a = self.make_tensor_abc_for_tensormap_update(
                        cur_group_idx,
                        ab_dtype,
                        (
                            problem_shape_m,
                            problem_shape_n,
                            problem_shape_k,
                        ),
                        ptrs_abc_smem,
                        0,  # 0 for tensor A
                    )
                    real_tensor_b = self.make_tensor_abc_for_tensormap_update(
                        cur_group_idx,
                        ab_dtype,
                        (
                            problem_shape_m,
                            problem_shape_n,
                            problem_shape_k,
                        ),
                        ptrs_abc_smem,
                        1,  # 1 for tensor B
                    )
                    real_tensor_sfa = self.make_tensor_sfasfb_for_tensormap_update(
                        cur_group_idx,
                        sf_dtype,
                        (
                            problem_shape_m,
                            problem_shape_n,
                            problem_shape_k,
                        ),
                        ptrs_sfasfb_smem,
                        0,  # 0 for tensor SFA
                    )
                    real_tensor_sfb = self.make_tensor_sfasfb_for_tensormap_update(
                        cur_group_idx,
                        sf_dtype,
                        (
                            problem_shape_m,
                            problem_shape_n,
                            problem_shape_k,
                        ),
                        ptrs_sfasfb_smem,
                        1,  # 1 for tensor SFB
                    )
                    if tensormap_init_done == False:
                        # wait tensormap initialization complete
                        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,
                        ),
                    )

                # M, N, L coordinates of the MMA tile within the current CTA tile
                mma_tile_coord_mnl = (
                    cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
                    cta_tile_idx_n,
                    0,
                )

                #
                # Slice to per mma tile index
                #
                # ((atom_v, rest_v), RestK)
                tAgA_slice = tAgA[
                    (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                ]
                # ((atom_v, rest_v), RestK)
                tBgB_slice = tBgB[
                    (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
                ]

                # ((atom_v, rest_v), RestK)
                tAgSFA_slice = tAgSFA[
                    (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                ]
                # ((atom_v, rest_v), RestK)
                tBgSFB_slice = tBgSFB[
                    (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
                ]

                # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
                #
                # Why: If you skip the peek and just call producer_acquire, you may stall immediately even when there’s no work ready to consume yet.
                # The peek lets the warp prefetch optimistically and only block if the stage is actually full. Without it, you can add unnecessary bubbles in the TMA pipeline.
                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)
                    tma_desc_ptr_a = tensormap_manager.get_tensormap_ptr(
                        tensormap_a_gmem_ptr,
                        cute.AddressSpace.generic,
                    )
                    tma_desc_ptr_b = tensormap_manager.get_tensormap_ptr(
                        tensormap_b_gmem_ptr,
                        cute.AddressSpace.generic,
                    )
                    tma_desc_ptr_sfa = tensormap_manager.get_tensormap_ptr(
                        tensormap_sfa_gmem_ptr,
                        cute.AddressSpace.generic,
                    )
                    tma_desc_ptr_sfb = tensormap_manager.get_tensormap_ptr(
                        tensormap_sfb_gmem_ptr,
                        cute.AddressSpace.generic,
                    )
                #
                # Tma load loop
                #
                for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
                    # Conditionally wait for AB buffer empty
                    ab_pipeline.producer_acquire(
                        ab_producer_state, peek_ab_empty_status
                    )

                    # Issue TMA load A/B/SFA/SFB
                    cute.copy(
                        tma_atom_a,
                        tAgA_slice[(None, ab_producer_state.count)],
                        tAsA[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=a_full_mcast_mask,
                        tma_desc_ptr=tma_desc_ptr_a,
                    )
                    cute.copy(
                        tma_atom_b,
                        tBgB_slice[(None, ab_producer_state.count)],
                        tBsB[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=b_full_mcast_mask,
                        tma_desc_ptr=tma_desc_ptr_b,
                    )
                    cute.copy(
                        tma_atom_sfa,
                        tAgSFA_slice[(None, ab_producer_state.count)],
                        tAsSFA[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=sfa_full_mcast_mask,
                        tma_desc_ptr=tma_desc_ptr_sfa,
                    )
                    cute.copy(
                        tma_atom_sfb,
                        tBgSFB_slice[(None, ab_producer_state.count)],
                        tBsSFB[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=sfb_full_mcast_mask,
                        tma_desc_ptr=tma_desc_ptr_sfb,
                    )

                    # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
                    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
                        )

                #
                # Advance to next tile
                #
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()
                last_group_idx = cur_group_idx

            #
            # Wait A/B buffer empty
            #
            ab_pipeline.producer_tail(ab_producer_state)

        #
        # Specialized MMA warp
        #
        if warp_idx == self.mma_warp_id:
            #
            # Initialize tensormaps for A, B, SFA and SFB
            #
            # Not specifically used in the MMA warp, one of the warp just had to initialise it and create
            # the base descriptior in the shared memory.
            #
            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
            )
            # indicate tensormap initialization has finished
            self.tensormap_ab_init_barrier.arrive_and_wait()

            #
            # Bar sync for retrieve tensor memory ptr from shared mem
            #
            self.tmem_alloc_barrier.arrive_and_wait()

            #
            # Retrieving tensor memory ptr and make accumulator/SFA/SFB tensor
            #
            # Make accumulator tmem tensor
            acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
                self.acc_dtype,
                alignment=16,
                ptr_to_buffer_holding_addr=tmem_holding_buf,
            )
            # (MMA, MMA_M, MMA_N, STAGE)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

            # Make SFA tmem tensor
            sfa_tmem_ptr = cute.recast_ptr(
                acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base),
                dtype=sf_dtype,
            )
            # (MMA, MMA_M, MMA_K)
            tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
                tiled_mma,
                self.mma_tiler,
                sf_vec_size,
                cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
            )
            tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)

            # Make SFB tmem tensor
            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=sf_dtype,
            )
            # (MMA, MMA_N, MMA_K)
            tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
                tiled_mma,
                self.mma_tiler,
                sf_vec_size,
                cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
            )
            tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
            #
            # Partition for S2T copy of SFA/SFB
            #
            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)
            )

            #
            # Persistent tile scheduling loop
            #
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), grid_dim
            )

            work_tile = tile_sched.initial_work_tile_info()
            mma_search_group_idx = cutlass.Int32(0)
            mma_search_tile_start = cutlass.Int32(0)
            mma_search_tile_end = group_tile_prefix_end[mma_search_group_idx]
            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:
                cur_tile_coord = work_tile.tile_idx
                # MMA warp is only interested in number of K tiles and group index.
                cur_linear_idx = cur_tile_coord[2]
                while (
                    cur_linear_idx >= mma_search_tile_end
                    and mma_search_group_idx + 1 < group_count
                ):
                    mma_search_group_idx += 1
                    mma_search_tile_start = mma_search_tile_end
                    mma_search_tile_end = group_tile_prefix_end[mma_search_group_idx]
                cur_group_idx = mma_search_group_idx
                cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]

                # (MMA, MMA_M, MMA_N)
                tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]

                # Peek (try_wait) AB buffer full for k_tile = 0
                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
                    )

                #
                # Wait for accumulator buffer empty
                #
                if is_leader_cta:
                    acc_pipeline.producer_acquire(acc_producer_state)

                #
                # Reset the ACCUMULATE field for each tile
                #
                tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

                #
                # Mma mainloop
                #
                for k_tile in range(cur_k_tile_cnt):
                    if is_leader_cta:
                        # Conditionally wait for AB buffer full
                        ab_pipeline.consumer_wait(
                            ab_consumer_state, peek_ab_full_status
                        )

                        #  Copy SFA/SFB from smem to tmem
                        s2t_stage_coord = (
                            None,
                            None,
                            None,
                            None,
                            ab_consumer_state.index,
                        )
                        tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
                        tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
                        cute.copy(
                            tiled_copy_s2t_sfa,
                            tCsSFA_compact_s2t_staged,
                            tCtSFA_compact_s2t,
                        )
                        cute.copy(
                            tiled_copy_s2t_sfb,
                            tCsSFB_compact_s2t_staged,
                            tCtSFB_compact_s2t,
                        )

                        # tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB
                        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,
                            )

                            # Set SFA/SFB tensor to tiled_mma
                            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,
                            )
                            # Enable accumulate after the first kblock only.
                            if kblock_idx == 0:
                                tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

                        # Async arrive AB buffer empty
                        ab_pipeline.consumer_release(ab_consumer_state)

                    # Peek (try_wait) AB buffer full for k_tile = k_tile + 1
                    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
                            )

                #
                # Async arrive accumulator buffer full
                #
                if is_leader_cta:
                    acc_pipeline.producer_commit(acc_producer_state)
                acc_producer_state.advance()

                #
                # Advance to next tile
                #
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()

            #
            # Wait for accumulator buffer empty
            #
            acc_pipeline.producer_tail(acc_producer_state)


        #
        # Specialized epilogue warps
        #
        if warp_idx < self.mma_warp_id:
            # initialize tensorap for C
            tensormap_manager.init_tensormap_from_atom(
                tma_atom_c,
                tensormap_c_smem_ptr,
                self.epilog_warp_id[0],
            )
            #
            # Alloc tensor memory buffer
            #
            if warp_idx == self.epilog_warp_id[0]:
                cute.arch.alloc_tmem(
                    num_tmem_alloc_cols,
                    tmem_holding_buf,
                    is_two_cta=use_2cta_instrs,
                )

            #
            # Bar sync for retrieve tensor memory ptr from shared memory
            #
            self.tmem_alloc_barrier.arrive_and_wait()

            #
            # Retrieving tensor memory ptr and make accumulator tensor
            #
            acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
                self.acc_dtype,
                alignment=16,
                ptr_to_buffer_holding_addr=tmem_holding_buf,
            )
            # (MMA, MMA_M, MMA_N, STAGE)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

            ### Start from here
            #
            # Partition for epilogue
            #
            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, 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
                )
            )

            #
            # Persistent tile scheduling loop
            #
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), grid_dim
            )

            work_tile = tile_sched.initial_work_tile_info()
            epi_search_group_idx = cutlass.Int32(0)
            epi_search_tile_start = cutlass.Int32(0)
            epi_search_tile_end = group_tile_prefix_end[epi_search_group_idx]

            acc_consumer_state = pipeline.make_pipeline_state(
                pipeline.PipelineUserType.Consumer, self.num_acc_stage
            )

            # Threads/warps participating in tma store pipeline
            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,
            )
            # group index to start searching
            last_group_idx = cutlass.Int32(-1)
            tma_desc_ptr_c = tensormap_manager.get_tensormap_ptr(
                tensormap_c_gmem_ptr,
                cute.AddressSpace.generic,
            )

            while work_tile.is_valid_tile:
                cur_tile_coord = work_tile.tile_idx
                cur_linear_idx = cur_tile_coord[2]
                while (
                    cur_linear_idx >= epi_search_tile_end
                    and epi_search_group_idx + 1 < group_count
                ):
                    epi_search_group_idx += 1
                    epi_search_tile_start = epi_search_tile_end
                    epi_search_tile_end = group_tile_prefix_end[epi_search_group_idx]
                cur_group_idx = epi_search_group_idx
                cluster_tile_idx = cur_linear_idx - epi_search_tile_start
                cluster_count_m = group_cluster_count_m[cur_group_idx]
                cta_tile_idx_m = (
                    (cluster_tile_idx % cluster_count_m) * self.cluster_shape_mn[0]
                    + cur_tile_coord[0]
                )
                cta_tile_idx_n = (
                    (cluster_tile_idx // cluster_count_m) * self.cluster_shape_mn[1]
                    + cur_tile_coord[1]
                )
                problem_shape_m = problem_sizes_smem[(cur_group_idx, 0)]
                problem_shape_n = problem_sizes_smem[(cur_group_idx, 1)]
                problem_shape_k = problem_sizes_smem[(cur_group_idx, 2)]
                is_group_changed = cur_group_idx != last_group_idx

                if is_group_changed:
                    # construct tensor c based on real shape, stride information
                    real_tensor_c = self.make_tensor_abc_for_tensormap_update(
                        cur_group_idx,
                        c_dtype,
                        (
                            problem_shape_m,
                            problem_shape_n,
                            problem_shape_k,
                        ),
                        ptrs_abc_smem,
                        2,  # 2 for tensor C
                    )
                    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 = (
                    cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
                    cta_tile_idx_n,
                    0,
                )
                cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]

                #
                # Slice to per mma tile index
                #
                # ((ATOM_V, REST_V), EPI_M, EPI_N)
                bSG_gC = bSG_gC_partitioned[
                    (
                        None,
                        None,
                        None,
                        *mma_tile_coord_mnl,
                    )
                ]

                # Set tensor memory buffer for current tile
                # (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
                tTR_tAcc = tTR_tAcc_base[
                    (None, None, None, None, None, acc_consumer_state.index)
                ]

                #
                # Wait for accumulator buffer full
                #
                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)
                    tma_desc_ptr_c = tensormap_manager.get_tensormap_ptr(
                        tensormap_c_gmem_ptr,
                        cute.AddressSpace.generic,
                    )

                #
                # Store accumulator to global memory in subtiles
                #
                subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
                num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
                tRS_rAcc = tiled_copy_r2s.retile(tTR_rAcc)
                for subtile_idx in range(subtile_cnt):
                    #
                    # Load accumulator from tensor memory buffer to register
                    #
                    tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
                    cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

                    #
                    # Convert to C type
                    #
                    acc_vec = tRS_rAcc.load()
                    tRS_rC.store(acc_vec.to(c_dtype))

                    #
                    # Store C to shared memory
                    #
                    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)],
                    )
                    # Fence and barrier to make sure shared memory store is visible to TMA store
                    cute.arch.fence_proxy(
                        cute.arch.ProxyKind.async_shared,
                        space=cute.arch.SharedSpace.shared_cta,
                    )
                    self.epilog_sync_barrier.arrive_and_wait()

                    #
                    # TMA store C to global memory
                    #
                    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=tma_desc_ptr_c,
                        )
                        # Fence and barrier to make sure shared memory store is visible to TMA store
                        c_pipeline.producer_commit()
                        c_pipeline.producer_acquire()
                    self.epilog_sync_barrier.arrive_and_wait()
                #
                # Async arrive accumulator buffer empty
                #
                with cute.arch.elect_one():
                    acc_pipeline.consumer_release(acc_consumer_state)
                acc_consumer_state.advance()

                #
                # Advance to next tile
                #
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()
                last_group_idx = cur_group_idx


            #
            # Dealloc the tensor memory buffer
            #
            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, num_tmem_alloc_cols, is_two_cta=use_2cta_instrs
                )
            #
            # Wait for C store complete
            #
            c_pipeline.producer_tail()



    @cute.jit
    def __call__(
        self,
        ptr_of_tensor_of_problem_sizes: cute.Pointer,
        ptr_of_tensor_of_abc_ptrs: cute.Pointer,
        ptr_of_tensor_of_sfasfb_ptrs: cute.Pointer,
        ptr_of_tensor_of_tensormap: cute.Pointer,
        sm_count: cutlass.Constexpr[int],
        total_num_clusters: cutlass.Constexpr[int],
        num_groups: cutlass.Constexpr[int],
        max_active_clusters: cutlass.Constexpr[int],
    ):
        tensor_of_abc_ptrs = cute.make_tensor(
            ptr_of_tensor_of_abc_ptrs, cute.make_layout((num_groups, 3), stride=(3, 1))
        )
        tensor_of_sfasfb_ptrs = cute.make_tensor(
            ptr_of_tensor_of_sfasfb_ptrs, cute.make_layout((num_groups, 2), stride=(2, 1))
        )
        problem_shape_mnkl = cute.make_tensor(
            ptr_of_tensor_of_problem_sizes, cute.make_layout((num_groups, 4), stride=(4, 1))
        )
        tensormap_stride1 = Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
        tensormap_stride0 = Sm100GroupedBlockScaledGemmKernel.num_tensormaps * tensormap_stride1
        tensor_of_tensormap = cute.make_tensor(
            ptr_of_tensor_of_tensormap,
            cute.make_layout(
                (sm_count, Sm100GroupedBlockScaledGemmKernel.num_tensormaps, tensormap_stride1),
                stride=(tensormap_stride0, tensormap_stride1, 1),
            ),
        )
        # Use fake shape for initial Tma descriptor and atom setup
        # The real Tma desc and atom will be updated during kernel execution.
        min_a_shape = (cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(1))
        min_b_shape = (cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(1))
        min_c_shape = (cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(1))
        initial_a = cute.make_tensor(
            cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16,),
            cute.make_layout(
                (min_a_shape[0], cute.assume(min_a_shape[2], 32), min_a_shape[3]),
                stride=(
                    cute.assume(min_a_shape[2], 32),
                    1,
                    cute.assume(min_a_shape[0] * min_a_shape[2], 32),
                ),
            ),
        )
        initial_b = cute.make_tensor(
            cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16,),
            cute.make_layout(
                (min_b_shape[1], cute.assume(min_b_shape[2], 32), min_b_shape[3]),
                stride=(
                    cute.assume(min_b_shape[2], 32),
                    1,
                    cute.assume(min_b_shape[1] * min_b_shape[2], 32),
                ),
            ),
        )
        initial_c = cute.make_tensor(
            cute.make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16,),
            cute.make_layout(
                (min_c_shape[1], cute.assume(min_c_shape[2], 32), min_c_shape[3]),
                stride=(
                    cute.assume(min_c_shape[2], 32),
                    1,
                    cute.assume(min_c_shape[1] * min_c_shape[2], 32),
                ),
            ),
        )

        # Setup sfa/sfb tensor by filling A/B tensor to scale factor atom layout
        # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL)
        sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
            initial_a.shape, sf_vec_size
        )
        # ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL)
        sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
            initial_b.shape, sf_vec_size
        )
        # Create initial SFA and SFB tensors with fake shape and null pointer.
        initial_sfa = cute.make_tensor(
            cute.make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16,), sfa_layout)
        initial_sfb = cute.make_tensor(
            cute.make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16,), sfb_layout)

        a_major_mode = utils.LayoutEnum.from_tensor(initial_a).mma_major_mode()
        b_major_mode = utils.LayoutEnum.from_tensor(initial_b).mma_major_mode()
        self.c_layout = utils.LayoutEnum.from_tensor(initial_c)

        mma_inst_shape_mn = (
            self.mma_tiler[0],
            self.mma_tiler[1],
        )

        mma_inst_shape_mn_sfb = (
            mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
            cute.round_up(mma_inst_shape_mn[1], 128),
        )


        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            ab_dtype,
            a_major_mode,
            b_major_mode,
            sf_dtype,
            sf_vec_size,
            self.cta_group,
            mma_inst_shape_mn,
        )

        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            ab_dtype,
            a_major_mode,
            b_major_mode,
            sf_dtype,
            sf_vec_size,
            cute.nvgpu.tcgen05.CtaGroup.ONE,
            mma_inst_shape_mn_sfb,
        )
        atom_thr_size = cute.size(tiled_mma.thr_id.shape)

        # Compute mma instruction shapes
        # (MMA_Tile_Shape_M, MMA_Tile_Shape_N, MMA_Inst_Shape_K)
        self.mma_inst_shape_mn = (
            self.mma_tiler[0],
            self.mma_tiler[1],
        )
        # (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K)
        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),
        )

        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))
        )

        # **Why* we compute the shape and what it's used for*
        # Because the epilogue needs a concrete tile shape to safely and efficiently write results back to global memory.
        # The MMA tile size and the optimal store tile size aren’t always the same—store needs to respect C’s layout, alignment, and vectorization/TMA requirements.
        # The “epilogue tile” tells the kernel how to break the accumulator into chunks, how big sC should be in shared memory,
        # and how the TMA store should be issued without out‑of‑bounds or misaligned accesses.
        epi_tile = sm100_utils.compute_epilogue_tile_shape(
            self.cta_tile_shape_mnk,
            self.use_2cta_instrs,
            self.c_layout,
            c_dtype,
        )

        # Stage counts are computed to maximize pipelining without exceeding SMEM: we estimate SMEM per A/B/SF stage and per C stage,
        # then choose the largest num_ab_stage that fits (given occupancy and reserved bytes), and use any leftover SMEM to add more C stages.
        # This balances latency hiding with SMEM limits and avoids over-allocation.
        self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
            tiled_mma,
            self.mma_tiler,
            ab_dtype,
            ab_dtype,
            epi_tile,
            c_dtype,
            self.c_layout,
            sf_dtype,
            sf_vec_size,
            smem_capacity,
            self.occupancy,
        )

        # Compute A/B/SFA/SFB/C shared memory layout
        # Because the SMEM layout for A is staged—it includes an extra “stage” dimension for pipelining (double/triple buffering).
        # The number of stages tells the layout how many A tiles to allocate in shared memory and how to index them.
        # So num_ab_stage directly determines the SMEM footprint and addressing for A (and similarly B/SF)
        self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(
            tiled_mma,
            self.mma_tiler,
            ab_dtype,
            self.num_ab_stage,
        )
        self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(
            tiled_mma,
            self.mma_tiler,
            ab_dtype,
            self.num_ab_stage,
        )
        self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            self.mma_tiler,
            sf_vec_size,
            self.num_ab_stage,
        )
        self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            self.mma_tiler,
            sf_vec_size,
            self.num_ab_stage,
        )
        self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
            c_dtype,
            self.c_layout,
            epi_tile,
            self.num_c_stage,
        )

        # Define shared storage for kernel
        @cute.struct
        class SharedStorage:
            # cute.struct.MemRange[T, N] is a typed shared‑memory buffer of N elements of type T, not bytes.
            # The size N is a count of elements. Total bytes = N * sizeof(T).
            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]
            problem_sizes_buffer: cute.struct.MemRange[cutlass.Int32, num_groups * 4]
            ptrs_abc_buffer: cute.struct.MemRange[cutlass.Int64, num_groups * 3]
            ptrs_sfasfb_buffer: cute.struct.MemRange[cutlass.Int64, num_groups * 2]
            group_tile_prefix_end: cute.struct.MemRange[cutlass.Int32, num_groups]
            group_cluster_count_m: cute.struct.MemRange[cutlass.Int32, num_groups]
            group_cluster_count_k: cute.struct.MemRange[cutlass.Int32, num_groups]

            # tmem_dealloc_mbar_ptr is the shared‑memory barrier (mbarrier) used to synchronize TMEM deallocation across warps/CTAs.
            # It ensures everyone is done using tensor memory before it’s freed.
            # It’s a single 64‑bit slot because mbarriers are represented as 64‑bit values in SMEM.
            tmem_dealloc_mbar_ptr: cutlass.Int64

            # tmem_holding_buf is a small shared‑memory slot used to hold the pointer to tensor memory (TMEM) that’s allocated at runtime.

            tmem_holding_buf: cutlass.Int32
            # (EPI_TILE_M, EPI_TILE_N, STAGE)
            sC: cute.struct.Align[
                cute.struct.MemRange[
                    c_dtype,
                    cute.cosize(self.c_smem_layout_staged.outer),
                ],
                buffer_align_bytes,
            ]
            # (MMA, MMA_M, MMA_K, STAGE)
            sA: cute.struct.Align[
                cute.struct.MemRange[
                    ab_dtype, cute.cosize(self.a_smem_layout_staged.outer)
                ],
                buffer_align_bytes,
            ]
            # (MMA, MMA_N, MMA_K, STAGE)
            sB: cute.struct.Align[
                cute.struct.MemRange[
                    ab_dtype, cute.cosize(self.b_smem_layout_staged.outer)
                ],
                buffer_align_bytes,
            ]
            # (MMA, MMA_M, MMA_K, STAGE)
            sSFA: cute.struct.Align[
                cute.struct.MemRange[
                    sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
                ],
                buffer_align_bytes,
            ]
            # (MMA, MMA_N, MMA_K, STAGE)
            sSFB: cute.struct.Align[
                cute.struct.MemRange[
                    sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
                ],
                buffer_align_bytes,
            ]

        self.shared_storage = SharedStorage

        # Compute cluster layout
        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,),
        )

        # Compute number of multicast CTAs for A/B
        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

        # Setup TMA load for A
        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,
        )

        # Setup TMA load for B
        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,
        )

        # Setup TMA load for SFA
        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,
        )

        # Setup TMA load for SFB
        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,
        )

        a_copy_size = cute.size_in_bytes(ab_dtype, a_smem_layout)
        b_copy_size = cute.size_in_bytes(ab_dtype, b_smem_layout)
        sfa_copy_size = cute.size_in_bytes(sf_dtype, sfa_smem_layout)
        sfb_copy_size = cute.size_in_bytes(sf_dtype, sfb_smem_layout)
        self.num_tma_load_bytes = (
            a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
        ) * atom_thr_size

        # Setup TMA store for C
        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,
            epi_tile,
        )

        # Compute grid size (For first problem)
        #   total_num_clusters = 352
        #   self.cluster_shape_mn = (1, 2)
        #   max_active_clusters = 148
        self.tile_sched_params, grid = self._compute_grid(
            total_num_clusters, self.cluster_shape_mn, max_active_clusters
        )
        # cute.printf("launch kernel\n")
        # cute.printf("grid: {}\n", grid),
        # cute.printf("block: [{}, 1, 1]\n", self.threads_per_cta)
        # cute.printf("cluster: {}\n", self.cluster_shape_mn)
        # cute.printf("num ab stage: {}\n", self.num_ab_stage)
        # cute.printf("num acc stage: {}\n", self.num_acc_stage)
        # cute.printf("num c stage: {}\n", self.num_c_stage)

        # Launch the kernel synchronously
        self.ported_kernel(
            tiled_mma=tiled_mma,
            tiled_mma_sfb=tiled_mma_sfb,
            tma_atom_a=tma_atom_a,
            mA_mkl=tma_tensor_a,
            tma_atom_b=tma_atom_b,
            mB_nkl=tma_tensor_b,
            tma_atom_sfa=tma_atom_sfa,
            mSFA_mkl=tma_tensor_sfa,
            tma_atom_sfb=tma_atom_sfb,
            mSFB_nkl=tma_tensor_sfb,
            tma_atom_c=tma_atom_c,
            mC_mnl=tma_tensor_c,
            cluster_layout_vmnk=self.cluster_layout_vmnk,
            cluster_layout_sfb_vmnk=self.cluster_layout_sfb_vmnk,
            a_smem_layout_staged=self.a_smem_layout_staged,
            b_smem_layout_staged=self.b_smem_layout_staged,
            sfa_smem_layout_staged=self.sfa_smem_layout_staged,
            sfb_smem_layout_staged=self.sfb_smem_layout_staged,
            c_smem_layout_staged=self.c_smem_layout_staged,
            epi_tile=epi_tile,
            tile_sched_params=self.tile_sched_params,
            group_count=num_groups,
            problem_sizes_mnkl=problem_shape_mnkl,
            ptrs_abc=tensor_of_abc_ptrs,
            ptrs_sfasfb=tensor_of_sfasfb_ptrs,
            tensormaps=tensor_of_tensormap,
        ).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.jit
    def make_tensor_abc_for_tensormap_update(
        self,
        group_idx: cutlass.Int32,
        dtype: Type[cutlass.Numeric],
        problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32],
        tensor_address_abc: cute.Tensor,
        tensor_index: int,
    ):
        """Construct a global tensor view for A, B or C for a specific group.

        This function is used within the kernel to dynamically create a CUTE tensor
        representing A, B, or C for the current group being processed, using the
        group-specific address and shape information.

        :param group_idx: The index of the current group within the grouped GEMM.
        :type group_idx: cutlass.Int32
        :param dtype: The data type of the tensor elements (e.g., cutlass.Float16).
        :type dtype: Type[cutlass.Numeric]
        :param problem_shape_mnk: The (M, N, K) problem shape for the current group.
        :type problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32]
        :param tensor_address_abc: Tensor containing global memory addresses for A, B, C for all groups. Layout: (group_count, 3).
        :type tensor_address_abc: cute.Tensor
        :param tensor_index: Specifies which tensor to create: 0 for A, 1 for B, 2 for C.
        :type tensor_index: int
        :return: A CUTE tensor representing the requested global memory tensor (A, B, or C) for the specified group.
        :rtype: cute.Tensor
        :raises TypeError: If the provided dtype is not a subclass of cutlass.Numeric.
        """
        ptr_i64 = tensor_address_abc[(group_idx, tensor_index)]
        if cutlass.const_expr(
            not isclass(dtype) or not issubclass(dtype, cutlass.Numeric)
        ):
            raise TypeError(
                f"dtype must be a type of cutlass.Numeric, got {type(dtype)}"
            )
        tensor_gmem_ptr = cute.make_ptr(
            dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16
        )

        c1 = cutlass.Int32(1)
        c0 = cutlass.Int32(0)

        if cutlass.const_expr(tensor_index == 0):  # tensor A
            m = problem_shape_mnk[0]
            k = problem_shape_mnk[2]
            return cute.make_tensor(
                tensor_gmem_ptr,
                cute.make_layout((m, k, c1), stride=(k, c1, c0)),
            )
        elif cutlass.const_expr(tensor_index == 1):  # tensor B
            n = problem_shape_mnk[1]
            k = problem_shape_mnk[2]
            return cute.make_tensor(
                tensor_gmem_ptr,
                cute.make_layout((n, k, c1), stride=(k, c1, c0)),
            )
        else:  # tensor C
            m = problem_shape_mnk[0]
            n = problem_shape_mnk[1]
            return cute.make_tensor(
                tensor_gmem_ptr,
                cute.make_layout((m, n, c1), stride=(n, c1, c0)),
            )

    @cute.jit
    def make_tensor_sfasfb_for_tensormap_update(
        self,
        group_idx: cutlass.Int32,
        dtype: Type[cutlass.Numeric],
        problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32],
        tensor_address_sfasfb: cute.Tensor,
        tensor_index: int,
    ):
        """Extract tensor address for a given group and construct a global tensor for SFA or SFB.

        This function is used within the kernel to dynamically create a CUTE tensor
        representing SFA or SFB for the current group being processed, using the
        group-specific address, shape information.

        :param group_idx: The index of the current group within the grouped GEMM.
        :type group_idx: cutlass.Int32
        :param dtype: The data type of the tensor elements (e.g., cutlass.Float16).
        :type dtype: Type[cutlass.Numeric]
        :param problem_shape_mnk: The (M, N, K) problem shape for the current group.
        :type problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32]
        :param tensor_address_sfasfb: Tensor containing global memory addresses for SFA, SFB for all groups. Layout: (group_count, 2).
        :type tensor_address_sfasfb: cute.Tensor
        :param tensor_index: Specifies which tensor to create: 0 for SFA, 1 for SFB.
        :type tensor_index: int
        :return: A CUTE tensor representing the requested global memory tensor (SFA, SFB) for the specified group.
        :rtype: cute.Tensor
        :raises TypeError: If the provided dtype is not a subclass of cutlass.Numeric.
        """
        ptr_i64 = tensor_address_sfasfb[(group_idx, tensor_index)]
        if cutlass.const_expr(
            not isclass(dtype) or not issubclass(dtype, cutlass.Numeric)
        ):
            raise TypeError(
                f"dtype must be a type of cutlass.Numeric, got {type(dtype)}"
            )
        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):  # tensor SFA
            m = problem_shape_mnk[0]
            k = problem_shape_mnk[2]
            sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
                (m, k, c1), sf_vec_size
            )
            return cute.make_tensor(
                tensor_gmem_ptr,
                sfa_layout,
            )
        else:  # tensor SFB
            n = problem_shape_mnk[1]
            k = problem_shape_mnk[2]
            sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
                (n, k, c1), sf_vec_size
            )
            return cute.make_tensor(
                tensor_gmem_ptr,
                sfb_layout,
            )

    @staticmethod
    def _compute_stages(
        tiled_mma: cute.TiledMma,
        mma_tiler_mnk: Tuple[int, int, int],
        a_dtype: Type[cutlass.Numeric],
        b_dtype: Type[cutlass.Numeric],
        epi_tile: cute.Tile,
        c_dtype: Type[cutlass.Numeric],
        c_layout: utils.LayoutEnum,
        sf_dtype: Type[cutlass.Numeric],
        sf_vec_size: int,
        smem_capacity: int,
        occupancy: int,
    ) -> Tuple[int, int, int]:
        """Computes the number of stages for A/B/C operands based on heuristics.

        :param tiled_mma: The tiled MMA object defining the core computation.
        :type tiled_mma: cute.TiledMma
        :param mma_tiler_mnk: The shape (M, N, K) of the MMA tiler.
        :type mma_tiler_mnk: tuple[int, int, int]
        :param a_dtype: Data type of operand A.
        :type a_dtype: type[cutlass.Numeric]
        :param b_dtype: Data type of operand B.
        :type b_dtype: type[cutlass.Numeric]
        :param epi_tile: The epilogue tile shape.
        :type epi_tile: cute.Tile
        :param c_dtype: Data type of operand C (output).
        :type c_dtype: type[cutlass.Numeric]
        :param c_layout: Layout enum of operand C.
        :type c_layout: utils.LayoutEnum
        :param sf_dtype: Data type of Scale factor.
        :type sf_dtype: type[cutlass.Numeric]
        :param sf_vec_size: Scale factor vector size.
        :type sf_vec_size: int
        :param smem_capacity: Total available shared memory capacity in bytes.
        :type smem_capacity: int
        :param occupancy: Target number of CTAs per SM (occupancy).
        :type occupancy: int

        :return: A tuple containing the computed number of stages for:
                 (ACC stages, A/B operand stages, C stages)
        :rtype: tuple[int, int, int]
        """
        # ACC stages
        num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2

        # Default C stages
        num_c_stage = 2

        # Calculate smem layout and size for one stage of A, B, SFA, SFB and C
        a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
            tiled_mma,
            mma_tiler_mnk,
            a_dtype,
            1,  # a tmp 1 stage is provided
        )
        b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
            tiled_mma,
            mma_tiler_mnk,
            b_dtype,
            1,  # a tmp 1 stage is provided
        )
        sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,  # a tmp 1 stage is provided
        )
        sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,  # a tmp 1 stage is provided
        )

        c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(
            c_dtype,
            c_layout,
            epi_tile,
            1,
        )

        ab_bytes_per_stage = (
            cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
            + cute.size_in_bytes(b_dtype, b_smem_layout_staged_one)
            + cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
            + cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
        )
        mbar_helpers_bytes = 1024
        c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one)
        c_bytes = c_bytes_per_stage * num_c_stage

        # Calculate A/B/SFA/SFB stages:
        # Start with total smem per CTA (capacity / occupancy)
        # Subtract reserved bytes and initial C stages bytes
        # Divide remaining by bytes needed per A/B/SFA/SFB stage
        num_ab_stage = (
            smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
        ) // ab_bytes_per_stage

        # Refine epilogue stages:
        # Calculate remaining smem after allocating for A/B/SFA/SFB stages and reserved bytes
        # Add remaining unused smem to epilogue
        num_c_stage += (
            smem_capacity
            - occupancy * ab_bytes_per_stage * num_ab_stage
            - occupancy * (mbar_helpers_bytes + c_bytes)
        ) // (occupancy * c_bytes_per_stage)

        # Optional manual tuning knobs.
        if max_num_ab_stage > 0:
            num_ab_stage = min(num_ab_stage, max_num_ab_stage)
        if force_num_c_stage > 0:
            num_c_stage = force_num_c_stage

        return num_acc_stage, num_ab_stage, num_c_stage

    @staticmethod
    def _compute_grid(
        total_num_clusters: int,
        cluster_shape_mn: tuple[int, int],
        max_active_clusters: cutlass.Constexpr[int],
    ) -> tuple[utils.PersistentTileSchedulerParams, tuple[int, int, int]]:
        """Compute tile scheduler parameters and grid shape for grouped GEMM operations.

        :param total_num_clusters: Total number of clusters to process across all groups.
        :type total_num_clusters: int
        :param cluster_shape_mn: Shape of each cluster in M, N dimensions.
        :type cluster_shape_mn: tuple[int, int]
        :param max_active_clusters: Maximum number of active clusters.
        :type max_active_clusters: cutlass.Constexpr[int]

        :return: A tuple containing:
            - tile_sched_params: Parameters for the persistent tile scheduler.
            - grid: Grid shape for kernel launch.
        :rtype: tuple[utils.PersistentTileSchedulerParams, tuple[int, ...]]
        """
        # Create problem shape with M, N dimensions from cluster shape
        # and L dimension representing the total number of clusters.
        problem_shape_ntile_mnl = (
            cluster_shape_mn[0],
            cluster_shape_mn[1],
            cutlass.Int32(total_num_clusters),
        )
        # (128, 128, 352)

        tile_sched_params = utils.PersistentTileSchedulerParams(
            problem_shape_ntile_mnl,
            (*cluster_shape_mn, 1),
            swizzle_size=tile_swizzle_size,
            raster_along_m=tile_raster_along_m,
        )

        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: int) -> int:
        """Calculate shared memory consumption for memory barriers based on provided stages.

        Each stage requires 2 barriers, and each barrier consumes 8 bytes of shared memory.
        The total consumption is the sum across all provided stages. This function calculates the total
        shared memory needed for these barriers.

        :param kwargs_stages: Variable keyword arguments where each key is a stage name
                              (e.g., num_acc_stage, num_ab_stage) and each value is the
                              number of stages of that type.
        :type kwargs_stages: int
        :return: Total shared memory bytes required for all memory barriers.
        :rtype: int
        """
        num_barriers_per_stage = 2
        num_bytes_per_barrier = 8
        mbar_smem_consumption = sum(
            [
                num_barriers_per_stage * num_bytes_per_barrier * stage
                for stage in kwargs_stages.values()
            ]
        )
        return mbar_smem_consumption

    def mainloop_s2t_copy_and_partition(
        self,
        sSF: cute.Tensor,
        tSF: cute.Tensor,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        """
        Make tiledCopy for smem to tmem load for scale factor tensor, then use it to partition smem memory (source) and tensor memory (destination).

        :param sSF: The scale factor tensor in smem
        :type sSF: cute.Tensor
        :param tSF: The scale factor tensor in tmem
        :type tSF: cute.Tensor

        :return: A tuple containing (tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t) where:
            - tiled_copy_s2t: The tiled copy operation for smem to tmem load for scale factor tensor(s2t)
            - tCsSF_compact_s2t: The partitioned scale factor tensor in smem
            - tSF_compact_s2t: The partitioned scale factor tensor in tmem
        :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]
        """
        # (MMA, MMA_MN, MMA_K, STAGE)
        tCsSF_compact = cute.filter_zeros(sSF)
        # (MMA, MMA_MN, MMA_K)
        tCtSF_compact = cute.filter_zeros(tSF)

        # Make S2T CopyAtom and tiledCopy
        copy_atom_s2t = cute.make_copy_atom(
            tcgen05.Cp4x32x128bOp(self.cta_group),
            sf_dtype,
        )
        tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
        thr_copy_s2t = tiled_copy_s2t.get_slice(0)

        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
        tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
        tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t, tCsSF_compact_s2t_
        )
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
        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],
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        """
        Make tiledCopy for tensor memory load, then use it to partition tensor memory (source) and register array (destination).

        :param tidx: The thread index in epilogue warp groups
        :type tidx: cutlass.Int32
        :param tAcc: The accumulator tensor to be copied and partitioned
        :type tAcc: cute.Tensor
        :param gC_mnl: The global tensor C
        :type gC_mnl: cute.Tensor
        :param epi_tile: The epilogue tiler
        :type epi_tile: cute.Tile
        :param use_2cta_instrs: Whether use_2cta_instrs is enabled
        :type use_2cta_instrs: bool

        :return: A tuple containing (tiled_copy_t2r, tTR_tAcc, tTR_rAcc) where:
            - tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r)
            - tTR_tAcc: The partitioned accumulator tensor
            - tTR_rAcc: The accumulated tensor in register used to hold t2r results
        :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]
        """
        # Make tiledCopy for tensor memory load
        copy_atom_t2r = sm100_utils.get_tmem_load_op(
            self.cta_tile_shape_mnk,
            self.c_layout,
            c_dtype,
            self.acc_dtype,
            epi_tile,
            use_2cta_instrs,
        )
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE)
        tAcc_epi = cute.flat_divide(
            tAcc[((None, None), 0, 0, None)],
            epi_tile,
        )
        # (EPI_TILE_M, EPI_TILE_N)
        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)
        # (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE)
        tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)

        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_mnl_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )
        # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        # (T2R, T2R_M, T2R_N)
        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,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        """
        Make tiledCopy for shared memory store, then use it to partition register array (source) and shared memory (destination).

        :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r)
        :type tiled_copy_t2r: cute.TiledCopy
        :param tTR_rC: The partitioned accumulator tensor
        :type tTR_rC: cute.Tensor
        :param tidx: The thread index in epilogue warp groups
        :type tidx: cutlass.Int32
        :param sC: The shared memory tensor to be copied and partitioned
        :type sC: cute.Tensor
        :type sepi: cute.Tensor

        :return: A tuple containing (tiled_copy_r2s, tRS_rC, tRS_sC) where:
            - tiled_copy_r2s: The tiled copy operation for register to smem copy(r2s)
            - tRS_rC: The partitioned tensor C (register source)
            - tRS_sC: The partitioned tensor C (smem destination)
        :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]
        """
        copy_atom_r2s = sm100_utils.get_smem_store_op(
            self.c_layout, c_dtype, self.acc_dtype, tiled_copy_t2r
        )
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
        # (R2S, R2S_M, R2S_N, PIPE_D)
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        # (R2S, R2S_M, R2S_N)
        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,
    ) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
        """Make tiledCopy for global memory store, then use it to:
        partition shared memory (source) and global memory (destination) for TMA store version.

        :param tidx: The thread index in epilogue warp groups
        :type tidx: cutlass.Int32
        :param atom: The copy_atom_c to be used for TMA store version, or tiled_copy_t2r for none TMA store version
        :type atom: cute.CopyAtom or cute.TiledCopy
        :param gC_mnl: The global tensor C
        :type gC_mnl: cute.Tensor
        :param epi_tile: The epilogue tiler
        :type epi_tile: cute.Tile
        :param sC: The shared memory tensor to be copied and partitioned
        :type sC: cute.Tensor

        :return: A tuple containing (tma_atom_c, bSG_sC, bSG_gC) where:
            - tma_atom_c: The TMA copy atom
            - bSG_sC: The partitioned shared memory tensor C
            - bSG_gC: The partitioned global tensor C
        :rtype: Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]
        """
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )

        tma_atom_c = atom
        sC_for_tma_partition = cute.group_modes(sC, 0, 2)
        gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
        # ((ATOM_V, REST_V), EPI_M, EPI_N)
        # ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
        bSG_sC, bSG_gC = cpasync.tma_partition(
            tma_atom_c,
            0,
            cute.make_layout(1),
            sC_for_tma_partition,
            gC_for_tma_partition,
        )
        return tma_atom_c, bSG_sC, bSG_gC



# Global caches to reduce repeated host-side setup overhead.
_compiled_kernel_cache = {}
_static_metadata_cache = {}
_tensormap_workspace_cache = {}
_pointer_metadata_cache = {}
_launch_context_cache = {}
_launch_context_fast_cache = {}
_hardware_info = None
_sm_count_cache = None
_max_active_clusters_cache = {}


def make_problem_sizes_cache_key(problem_sizes: List[Tuple[int, int, int, int]]) -> str:
    """Build a deterministic cache key from the full ordered problem-size tuples."""
    if not problem_sizes:
        return "g0"
    packed_groups = "|".join(f"{m},{n},{k},{l}" for m, n, k, l in problem_sizes)
    return f"g{len(problem_sizes)}:{packed_groups}"


def _get_hardware_info() -> cutlass.utils.HardwareInfo:
    global _hardware_info
    if _hardware_info is None:
        _hardware_info = cutlass.utils.HardwareInfo()
    return _hardware_info


def _get_sm_count() -> int:
    global _sm_count_cache
    if _sm_count_cache is None:
        _sm_count_cache = _get_hardware_info().get_max_active_clusters(1)
    return _sm_count_cache


def _get_max_active_clusters(cluster_shape_mn: Tuple[int, int]) -> int:
    cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]
    if cluster_size not in _max_active_clusters_cache:
        _max_active_clusters_cache[cluster_size] = (
            _get_hardware_info().get_max_active_clusters(cluster_size)
        )
    return _max_active_clusters_cache[cluster_size]


def _get_static_metadata_tensors(
    problem_sizes: List[Tuple[int, int, int, int]],
) -> torch.Tensor:
    device_idx = torch.cuda.current_device()
    cache_key = (device_idx, make_problem_sizes_cache_key(problem_sizes))
    cached = _static_metadata_cache.get(cache_key)
    if cached is not None:
        return cached

    # (num_groups, 4)
    tensor_of_problem_sizes = torch.tensor(
        problem_sizes, dtype=torch.int32, device="cuda"
    )
    _static_metadata_cache[cache_key] = tensor_of_problem_sizes
    return tensor_of_problem_sizes


def _get_pointer_metadata_tensors(
    abc_ptrs,
    sfasfb_ptrs,
) -> Tuple[torch.Tensor, torch.Tensor]:
    device_idx = torch.cuda.current_device()
    cache_key = (
        device_idx,
        tuple(abc_ptrs),
        tuple(sfasfb_ptrs),
    )
    cached = _pointer_metadata_cache.get(cache_key)
    if cached is not None:
        return cached

    tensor_of_abc_ptrs = torch.tensor(abc_ptrs, dtype=torch.int64, device="cuda")
    tensor_of_sfasfb_ptrs = torch.tensor(sfasfb_ptrs, dtype=torch.int64, device="cuda")
    _pointer_metadata_cache[cache_key] = (tensor_of_abc_ptrs, tensor_of_sfasfb_ptrs)
    return tensor_of_abc_ptrs, tensor_of_sfasfb_ptrs


def _get_tensormap_workspace(sm_count: int) -> torch.Tensor:
    device_idx = torch.cuda.current_device()
    cache_key = (device_idx, sm_count)
    tensormap_shape = (
        sm_count,
        Sm100GroupedBlockScaledGemmKernel.num_tensormaps,
        Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8,
    )
    cached = _tensormap_workspace_cache.get(cache_key)
    if cached is not None and tuple(cached.shape) == tensormap_shape:
        return cached

    tensor_of_tensormap = torch.empty(tensormap_shape, dtype=torch.int64, device="cuda")
    _tensormap_workspace_cache[cache_key] = tensor_of_tensormap
    return tensor_of_tensormap


# This function is used to compile the kernel once and cache it and then allow users to
# run the kernel multiple times to get more accurate timing results.
def compile_kernel(problem_sizes):
    """
    Compile the kernel once and cache it using problem_sizes as the key.
    This should be called before any timing measurements.

    Returns:
        The compiled kernel function
    """
    global _compiled_kernel_cache

    # Build cache key from the full ordered (m, n, k, l) tuples.
    cache_key = make_problem_sizes_cache_key(problem_sizes)

    # Check if we already have a compiled kernel for these problem sizes
    if cache_key in _compiled_kernel_cache:
        return _compiled_kernel_cache[cache_key]

    # Dummy/null pointers are only used to define the kernel ABI for compile-time;
    # real device pointers are passed at runtime from custom_kernel.
    cute_ptr_of_tensor_of_problem_sizes = make_ptr(
        cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16,
    )
    cute_ptr_of_tensor_of_abc_ptrs = make_ptr(
        cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16,
    )
    cute_ptr_of_tensor_of_sfasfb_ptrs = make_ptr(
        cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16,
    )

    # Each cluster needs its own set of tensormaps (one for A, B, SFA, SFB)
    # Shape: (total_num_clusters, num_tensormaps=5, bytes_per_tensormap/8=16)
    cute_ptr_of_tensor_of_tensormap = make_ptr(
        cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16,
    )

    kernel_runner = Sm100GroupedBlockScaledGemmKernel(
        (mma_tiler_mnk[0], mma_tiler_mnk[1]),
        cluster_shape_mn=cluster_shape_mn
    )


    # Fake cluster numbers for compile only.
    # Compute the tile shape for each CUDA Thread Block (CTA)
    # cta_tile_shape_mn: [M_tile, N_tile] = [128, 128] for this kernel
    cta_tile_shape_mn = [mma_tiler_mnk[0], mma_tiler_mnk[1]]
    # cluster_tile_shape_mn: Total tile shape per cluster
    cluster_tile_shape_mn = tuple(
        x * y for x, y in zip(cta_tile_shape_mn, cluster_shape_mn)
    )

    # Compute total number of cluster tiles needed across all groups
    # Each group's (m, n) dimensions are divided into tiles of size cluster_tile_shape_mn
    # This determines the total grid size (bidz dimension) for kernel launch
    total_num_clusters = 0
    for m, n, _, _ in problem_sizes:
        # Calculate number of tiles needed in M and N dimensions for this group
        num_clusters_mn = tuple(
            (x + y - 1) // y for x, y in zip((m, n), cluster_tile_shape_mn)
        )
        # Multiply M_tiles * N_tiles to get total tiles for this group
        total_num_clusters += functools.reduce(lambda x, y: x * y, num_clusters_mn)

    max_active_clusters = _get_max_active_clusters(cluster_shape_mn)
    sm_count = _get_sm_count()
    num_groups = len(problem_sizes)

    compiled_func = cute.compile(
        kernel_runner,
        cute_ptr_of_tensor_of_problem_sizes,
        cute_ptr_of_tensor_of_abc_ptrs,
        cute_ptr_of_tensor_of_sfasfb_ptrs,
        cute_ptr_of_tensor_of_tensormap,
        sm_count,
        total_num_clusters,
        num_groups,
        max_active_clusters
    )

    # Store compiled kernel in cache with problem_sizes as key
    _compiled_kernel_cache[cache_key] = compiled_func
    return compiled_func


def custom_kernel(data: input_t) -> output_t:
    """
    Execute the block-scaled group GEMM kernel.

    This is the main entry point called by the evaluation framework.
    It converts PyTorch tensors to CuTe tensors, launches the kernel,
    and returns the result.

    Args:
        data: Tuple of (abc_tensors, sfasfb_tensors, problem_sizes) where:
            abc_tensors: list of tuples (a, b, c) where
                a is torch.Tensor[float4e2m1fn_x2] of shape [m, k // 2, l]
                b is torch.Tensor[float4e2m1fn_x2] of shape [n, k // 2, l]
                c is torch.Tensor[float16] of shape [m, n, l]
            sfasfb_tensors: list of tuples (sfa, sfb) where
                sfa is torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l]
                sfb is torch.Tensor[float8_e4m3fnuz] of shape [n, k // 16, l]
            problem_sizes: list of tuples (m, n, k, l)
            each group has its own a, b, c, sfa, sfb with different m, n, k, l problem sizes
            l should always be 1 for each group.
            list size is the number of groups.

    Returns:
        list of c tensors where c is torch.Tensor[float16] of shape [m, n, l] for each group
    """
    abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
    device_idx = torch.cuda.current_device()
    problem_key = make_problem_sizes_cache_key(problem_sizes)

    # Fast path cache keyed by object identity. Guard against Python id reuse by
    # verifying referential identity before trusting the cached launch context.
    fast_cache_key = (
        device_idx,
        problem_key,
        id(abc_tensors),
        id(sfasfb_reordered_tensors),
    )
    cached_fast = _launch_context_fast_cache.get(fast_cache_key)
    if cached_fast is not None:
        cached_abc_obj, cached_sfa_obj, cached_launch = cached_fast
        if cached_abc_obj is not abc_tensors or cached_sfa_obj is not sfasfb_reordered_tensors:
            _launch_context_fast_cache.pop(fast_cache_key, None)
            cached_launch = None
    else:
        cached_launch = None

    if cached_launch is None:
        abc_ptrs = tuple((a.data_ptr(), b.data_ptr(), c.data_ptr()) for a, b, c in abc_tensors)
        sfasfb_ptrs = tuple(
            (sfa_reordered.data_ptr(), sfb_reordered.data_ptr())
            for sfa_reordered, sfb_reordered in sfasfb_reordered_tensors
        )
        launch_cache_key = (
            device_idx,
            problem_key,
            abc_ptrs,
            sfasfb_ptrs,
        )
        cached_launch = _launch_context_cache.get(launch_cache_key)
    if cached_launch is None:
        compiled_func = compile_kernel(problem_sizes)

        tensor_of_problem_sizes = _get_static_metadata_tensors(problem_sizes)

        # Cache pointer tensors to avoid per-call CPU work and H2D metadata upload.
        # tensor_of_abc_ptrs: Shape (num_groups, 3) containing (a_ptr, b_ptr, c_ptr) per group
        # tensor_of_sfasfb_ptrs: Shape (num_groups, 2) containing (sfa_ptr, sfb_ptr) per group
        tensor_of_abc_ptrs, tensor_of_sfasfb_ptrs = _get_pointer_metadata_tensors(
            abc_ptrs,
            sfasfb_ptrs,
        )

        # Allocate device memory for tensormap descriptors
        # Each cluster needs its own set of tensormaps (one for A, B, SFA, SFB, C)
        # Shape: (sm_count, num_tensormaps=5, bytes_per_tensormap/8=16)
        #   sm_count because we are running persisent scheduler.
        # Tensormaps are hardware descriptors used by TMA for efficient memory transfers
        sm_count = _get_sm_count()
        tensor_of_tensormap = _get_tensormap_workspace(sm_count)

        # Create CuTe pointers to the metadata tensors that will be passed to the kernel
        # These allow the GPU kernel to read problem sizes and tensor pointers
        cute_ptr_of_tensor_of_abc_ptrs = make_ptr(
            cutlass.Int64,
            tensor_of_abc_ptrs.data_ptr(),
            cute.AddressSpace.gmem,
            assumed_align=16,
        )
        cute_ptr_of_tensor_of_sfasfb_ptrs = make_ptr(
            cutlass.Int64,
            tensor_of_sfasfb_ptrs.data_ptr(),
            cute.AddressSpace.gmem,
            assumed_align=16,
        )
        cute_ptr_of_tensor_of_problem_sizes = make_ptr(
            cutlass.Int32,
            tensor_of_problem_sizes.data_ptr(),
            cute.AddressSpace.gmem,
            assumed_align=16,
        )
        cute_ptr_of_tensor_of_tensormap = make_ptr(
            cutlass.Int64,
            tensor_of_tensormap.data_ptr(),
            cute.AddressSpace.gmem,
            assumed_align=16,
        )
        cached_launch = (
            compiled_func,
            cute_ptr_of_tensor_of_problem_sizes,
            cute_ptr_of_tensor_of_abc_ptrs,
            cute_ptr_of_tensor_of_sfasfb_ptrs,
            cute_ptr_of_tensor_of_tensormap,
        )
        _launch_context_cache[launch_cache_key] = cached_launch
    _launch_context_fast_cache[fast_cache_key] = (
        abc_tensors,
        sfasfb_reordered_tensors,
        cached_launch,
    )
    (
        compiled_func,
        cute_ptr_of_tensor_of_problem_sizes,
        cute_ptr_of_tensor_of_abc_ptrs,
        cute_ptr_of_tensor_of_sfasfb_ptrs,
        cute_ptr_of_tensor_of_tensormap,
    ) = cached_launch

    # Launch the JIT-compiled GPU kernel with all prepared data
    # The kernel will perform block-scaled group GEMM: C = A * SFA * B * SFB for all groups
    compiled_func(
        cute_ptr_of_tensor_of_problem_sizes, # Pointer to problem sizes array
        cute_ptr_of_tensor_of_abc_ptrs,      # Pointer to ABC tensor pointers array
        cute_ptr_of_tensor_of_sfasfb_ptrs,   # Pointer to scale factor pointers array
        cute_ptr_of_tensor_of_tensormap,     # Pointer to tensormap buffer
    )

    num_groups = len(problem_sizes)
    res = []
    for i in range(num_groups):
        res.append(abc_tensors[i][2])
    return res
scrolls · 2574 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 480794.

⋯ 41 unchanged lines
# Total number of columns in tmem
num_tmem_alloc_cols = 512
+ # Staging knobs for tuning.
+ # max_num_ab_stage=0 means no cap.
+ # force_num_c_stage=0 means keep computed value.
+ max_num_ab_stage = 6
+ force_num_c_stage = 0
+
+ # Persistent scheduler knobs.
+ # tile_swizzle_size=1 means no swizzle.
+ tile_swizzle_size = 1
+ # Only used when tile_swizzle_size > 1.
+ tile_raster_along_m = True
+
smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
buffer_align_bytes = 1024
⋯ 13 unchanged lines
- and you’re targeting performance on SM100.
"""
- cluster_shape_mn = (1, 2)
+ cluster_shape_mn = (1, 1)
class Sm100GroupedBlockScaledGemmKernel:
# Size of smem we reserved for mbarrier, tensor memory management and tensormap update
⋯ 92 unchanged lines
ptrs_abc: cute.Tensor,
ptrs_sfasfb: cute.Tensor,
tensormaps: cute.Tensor,
- strides_abc: cute.Tensor
):
warp_idx = cute.arch.warp_idx()
# Hinting to the compiler this is identical across warp, enables warp level control flow, no divergent branches, better codegen.
⋯ 126 unchanged lines
# (MMA, MMA_N, MMA_K, STAGE)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
+ # Stage small per-group metadata into shared memory once per CTA to avoid
+ # repeatedly reloading it from global memory in persistent loops.
+ problem_sizes_smem = storage.problem_sizes_buffer.get_tensor(
+ cute.make_layout((group_count, 4), stride=(4, 1))
+ )
+ ptrs_abc_smem = storage.ptrs_abc_buffer.get_tensor(
+ cute.make_layout((group_count, 3), stride=(3, 1))
+ )
+ ptrs_sfasfb_smem = storage.ptrs_sfasfb_buffer.get_tensor(
+ cute.make_layout((group_count, 2), stride=(2, 1))
+ )
+ group_tile_prefix_end = storage.group_tile_prefix_end.get_tensor(
+ cute.make_layout(group_count, stride=1)
+ )
+ group_cluster_count_m = storage.group_cluster_count_m.get_tensor(
+ cute.make_layout(group_count, stride=1)
+ )
+ group_cluster_count_k = storage.group_cluster_count_k.get_tensor(
+ cute.make_layout(group_count, stride=1)
+ )
+ for linear_idx in cutlass.range(tidx, group_count * 4, self.threads_per_cta):
+ group_idx = linear_idx // 4
+ mode_idx = linear_idx % 4
+ problem_sizes_smem[(group_idx, mode_idx)] = problem_sizes_mnkl[
+ (group_idx, mode_idx)
+ ]
+ for linear_idx in cutlass.range(tidx, group_count * 3, self.threads_per_cta):
+ group_idx = linear_idx // 3
+ mode_idx = linear_idx % 3
+ ptrs_abc_smem[(group_idx, mode_idx)] = ptrs_abc[(group_idx, mode_idx)]
+ for linear_idx in cutlass.range(tidx, group_count * 2, self.threads_per_cta):
+ group_idx = linear_idx // 2
+ mode_idx = linear_idx % 2
+ ptrs_sfasfb_smem[(group_idx, mode_idx)] = ptrs_sfasfb[
+ (group_idx, mode_idx)
+ ]
+ if tidx == 0:
+ tile_prefix_end = cutlass.Int32(0)
+ for group_idx in cutlass.range(group_count):
+ problem_m = problem_sizes_smem[(group_idx, 0)]
+ problem_n = problem_sizes_smem[(group_idx, 1)]
+ problem_k = problem_sizes_smem[(group_idx, 2)]
+ cluster_count_m = (
+ problem_m + self.cluster_tile_shape_mnk[0] - 1
+ ) // self.cluster_tile_shape_mnk[0]
+ cluster_count_n = (
+ problem_n + self.cluster_tile_shape_mnk[1] - 1
+ ) // self.cluster_tile_shape_mnk[1]
+ cluster_count_k = (
+ problem_k + self.cluster_tile_shape_mnk[2] - 1
+ ) // self.cluster_tile_shape_mnk[2]
+ tile_prefix_end += cluster_count_m * cluster_count_n
+ group_tile_prefix_end[group_idx] = tile_prefix_end
+ group_cluster_count_m[group_idx] = cluster_count_m
+ group_cluster_count_k[group_idx] = cluster_count_k
+ cute.arch.barrier()
+
#
# Compute multicast mask for A/B/SFA/SFB buffer full
#
⋯ 200 unchanged lines
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), grid_dim
)
- # grouped gemm tile scheduler helper will compute the group index for the tile we're working on
- group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
- group_count,
- tile_sched_params,
- self.cluster_tile_shape_mnk,
- utils.create_initial_search_state(),
- )
tensormap_init_done = cutlass.Boolean(False)
# group index of last tile
last_group_idx = cutlass.Int32(-1)
+ tma_desc_ptr_a = tensormap_manager.get_tensormap_ptr(
+ tensormap_a_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_b = tensormap_manager.get_tensormap_ptr(
+ tensormap_b_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_sfa = tensormap_manager.get_tensormap_ptr(
+ tensormap_sfa_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_sfb = tensormap_manager.get_tensormap_ptr(
+ tensormap_sfb_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
work_tile = tile_sched.initial_work_tile_info()
+ tma_search_group_idx = cutlass.Int32(0)
+ tma_search_tile_start = cutlass.Int32(0)
+ tma_search_tile_end = group_tile_prefix_end[tma_search_group_idx]
# If producer and consumer use different num_stages, they’ll index different ring sizes and the barrier protocol breaks (stages will be marked full/empty out of sync).
# In practice, that would deadlock or corrupt data. The pipeline API assumes both sides are created with the same stage count.
⋯ 3 unchanged lines
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
- grouped_gemm_cta_tile_info = group_gemm_ts_helper.delinearize_z(
- cur_tile_coord,
- problem_sizes_mnkl,
+ cur_linear_idx = cur_tile_coord[2]
+ while (
+ cur_linear_idx >= tma_search_tile_end
+ and tma_search_group_idx + 1 < group_count
+ ):
+ tma_search_group_idx += 1
+ tma_search_tile_start = tma_search_tile_end
+ tma_search_tile_end = group_tile_prefix_end[tma_search_group_idx]
+ cur_group_idx = tma_search_group_idx
+ cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]
+ cluster_tile_idx = cur_linear_idx - tma_search_tile_start
+ cluster_count_m = group_cluster_count_m[cur_group_idx]
+ cta_tile_idx_m = (
+ (cluster_tile_idx % cluster_count_m) * self.cluster_shape_mn[0]
+ + cur_tile_coord[0]
)
- cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
- cur_group_idx = grouped_gemm_cta_tile_info.group_idx
+ cta_tile_idx_n = (
+ (cluster_tile_idx // cluster_count_m) * self.cluster_shape_mn[1]
+ + cur_tile_coord[1]
+ )
+ problem_shape_m = problem_sizes_smem[(cur_group_idx, 0)]
+ problem_shape_n = problem_sizes_smem[(cur_group_idx, 1)]
+ problem_shape_k = problem_sizes_smem[(cur_group_idx, 2)]
is_group_changed = cur_group_idx != last_group_idx
# skip tensormap update if we're working on the same group
if is_group_changed:
⋯ 1 unchanged lines
cur_group_idx,
ab_dtype,
(
- grouped_gemm_cta_tile_info.problem_shape_m,
- grouped_gemm_cta_tile_info.problem_shape_n,
- grouped_gemm_cta_tile_info.problem_shape_k,
+ problem_shape_m,
+ problem_shape_n,
+ problem_shape_k,
),
- strides_abc,
- ptrs_abc,
+ ptrs_abc_smem,
0, # 0 for tensor A
)
real_tensor_b = self.make_tensor_abc_for_tensormap_update(
cur_group_idx,
ab_dtype,
(
- grouped_gemm_cta_tile_info.problem_shape_m,
- grouped_gemm_cta_tile_info.problem_shape_n,
- grouped_gemm_cta_tile_info.problem_shape_k,
+ problem_shape_m,
+ problem_shape_n,
+ problem_shape_k,
),
- strides_abc,
- ptrs_abc,
+ ptrs_abc_smem,
1, # 1 for tensor B
)
real_tensor_sfa = self.make_tensor_sfasfb_for_tensormap_update(
cur_group_idx,
sf_dtype,
(
- grouped_gemm_cta_tile_info.problem_shape_m,
- grouped_gemm_cta_tile_info.problem_shape_n,
- grouped_gemm_cta_tile_info.problem_shape_k,
+ problem_shape_m,
+ problem_shape_n,
+ problem_shape_k,
),
- ptrs_sfasfb,
+ ptrs_sfasfb_smem,
0, # 0 for tensor SFA
)
real_tensor_sfb = self.make_tensor_sfasfb_for_tensormap_update(
cur_group_idx,
sf_dtype,
(
- grouped_gemm_cta_tile_info.problem_shape_m,
- grouped_gemm_cta_tile_info.problem_shape_n,
- grouped_gemm_cta_tile_info.problem_shape_k,
+ problem_shape_m,
+ problem_shape_n,
+ problem_shape_k,
),
- ptrs_sfasfb,
+ ptrs_sfasfb_smem,
1, # 1 for tensor SFB
)
if tensormap_init_done == False:
⋯ 26 unchanged lines
# M, N, L coordinates of the MMA tile within the current CTA tile
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,
+ cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
+ cta_tile_idx_n,
0,
)
⋯ 34 unchanged lines
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)
- tma_desc_ptr_a = tensormap_manager.get_tensormap_ptr(
- tensormap_a_gmem_ptr,
- cute.AddressSpace.generic,
- )
- tma_desc_ptr_b = tensormap_manager.get_tensormap_ptr(
- tensormap_b_gmem_ptr,
- cute.AddressSpace.generic,
- )
- tma_desc_ptr_sfa = tensormap_manager.get_tensormap_ptr(
- tensormap_sfa_gmem_ptr,
- cute.AddressSpace.generic,
- )
- tma_desc_ptr_sfb = tensormap_manager.get_tensormap_ptr(
- tensormap_sfb_gmem_ptr,
- cute.AddressSpace.generic,
- )
+ tma_desc_ptr_a = tensormap_manager.get_tensormap_ptr(
+ tensormap_a_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_b = tensormap_manager.get_tensormap_ptr(
+ tensormap_b_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_sfa = tensormap_manager.get_tensormap_ptr(
+ tensormap_sfa_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
+ tma_desc_ptr_sfb = tensormap_manager.get_tensormap_ptr(
+ tensormap_sfb_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
#
# Tma load loop
#
⋯ 144 unchanged lines
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), grid_dim
)
- # grouped gemm tile scheduler helper will compute the group index for the tile we're working on
- group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
- group_count,
- tile_sched_params,
- self.cluster_tile_shape_mnk,
- utils.create_initial_search_state(),
- )
work_tile = tile_sched.initial_work_tile_info()
+ mma_search_group_idx = cutlass.Int32(0)
+ mma_search_tile_start = cutlass.Int32(0)
+ mma_search_tile_end = group_tile_prefix_end[mma_search_group_idx]
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
⋯ 2 unchanged lines
)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
- # MMA warp is only interested in number of tiles along K dimension
- (
- cur_k_tile_cnt,
- cur_group_idx,
- ) = group_gemm_ts_helper.search_cluster_tile_count_k(
- cur_tile_coord,
- problem_sizes_mnkl,
- )
+ # MMA warp is only interested in number of K tiles and group index.
+ cur_linear_idx = cur_tile_coord[2]
+ while (
+ cur_linear_idx >= mma_search_tile_end
+ and mma_search_group_idx + 1 < group_count
+ ):
+ mma_search_group_idx += 1
+ mma_search_tile_start = mma_search_tile_end
+ mma_search_tile_end = group_tile_prefix_end[mma_search_group_idx]
+ cur_group_idx = mma_search_group_idx
+ cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]
# (MMA, MMA_M, MMA_N)
tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
⋯ 174 unchanged lines
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), grid_dim
)
- # grouped gemm tile scheduler helper will compute the group index for the tile we're working on
- group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
- group_count,
- tile_sched_params,
- self.cluster_tile_shape_mnk,
- utils.create_initial_search_state(),
- )
work_tile = tile_sched.initial_work_tile_info()
+ epi_search_group_idx = cutlass.Int32(0)
+ epi_search_tile_start = cutlass.Int32(0)
+ epi_search_tile_end = group_tile_prefix_end[epi_search_group_idx]
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
⋯ 10 unchanged lines
)
# group index to start searching
last_group_idx = cutlass.Int32(-1)
+ tma_desc_ptr_c = tensormap_manager.get_tensormap_ptr(
+ tensormap_c_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
- grouped_gemm_cta_tile_info = group_gemm_ts_helper.delinearize_z(
- cur_tile_coord,
- problem_sizes_mnkl,
+ cur_linear_idx = cur_tile_coord[2]
+ while (
+ cur_linear_idx >= epi_search_tile_end
+ and epi_search_group_idx + 1 < group_count
+ ):
+ epi_search_group_idx += 1
+ epi_search_tile_start = epi_search_tile_end
+ epi_search_tile_end = group_tile_prefix_end[epi_search_group_idx]
+ cur_group_idx = epi_search_group_idx
+ cluster_tile_idx = cur_linear_idx - epi_search_tile_start
+ cluster_count_m = group_cluster_count_m[cur_group_idx]
+ cta_tile_idx_m = (
+ (cluster_tile_idx % cluster_count_m) * self.cluster_shape_mn[0]
+ + cur_tile_coord[0]
)
- cur_group_idx = grouped_gemm_cta_tile_info.group_idx
+ cta_tile_idx_n = (
+ (cluster_tile_idx // cluster_count_m) * self.cluster_shape_mn[1]
+ + cur_tile_coord[1]
+ )
+ problem_shape_m = problem_sizes_smem[(cur_group_idx, 0)]
+ problem_shape_n = problem_sizes_smem[(cur_group_idx, 1)]
+ problem_shape_k = problem_sizes_smem[(cur_group_idx, 2)]
is_group_changed = cur_group_idx != last_group_idx
if is_group_changed:
⋯ 2 unchanged lines
cur_group_idx,
c_dtype,
(
- grouped_gemm_cta_tile_info.problem_shape_m,
- grouped_gemm_cta_tile_info.problem_shape_n,
- grouped_gemm_cta_tile_info.problem_shape_k,
+ problem_shape_m,
+ problem_shape_n,
+ problem_shape_k,
),
- strides_abc,
- ptrs_abc,
+ ptrs_abc_smem,
2, # 2 for tensor C
)
tensormap_manager.update_tensormap(
⋯ 5 unchanged lines
)
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,
+ cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
+ cta_tile_idx_n,
0,
)
- cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
+ cur_k_tile_cnt = group_cluster_count_k[cur_group_idx]
#
# Slice to per mma tile index
⋯ 25 unchanged lines
if is_group_changed:
if warp_idx == self.epilog_warp_id[0]:
tensormap_manager.fence_tensormap_update(tensormap_c_gmem_ptr)
+ tma_desc_ptr_c = tensormap_manager.get_tensormap_ptr(
+ tensormap_c_gmem_ptr,
+ cute.AddressSpace.generic,
+ )
#
# Store accumulator to global memory in subtiles
#
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
- tma_desc_ptr_c = tensormap_manager.get_tensormap_ptr(
- tensormap_c_gmem_ptr,
- cute.AddressSpace.generic,
- )
tRS_rAcc = tiled_copy_r2s.retile(tTR_rAcc)
for subtile_idx in range(subtile_cnt):
#
⋯ 83 unchanged lines
ptr_of_tensor_of_abc_ptrs: cute.Pointer,
ptr_of_tensor_of_sfasfb_ptrs: cute.Pointer,
ptr_of_tensor_of_tensormap: cute.Pointer,
- ptr_of_tensor_strides_abc: cute.Pointer,
sm_count: cutlass.Constexpr[int],
total_num_clusters: cutlass.Constexpr[int],
num_groups: cutlass.Constexpr[int],
⋯ 17 unchanged lines
stride=(tensormap_stride0, tensormap_stride1, 1),
),
)
- # provides 3D logical view onto 1D memory
- # - The pointer gives the base address of the flat data in memory.
- # - The layout (shape + stride) tells CuTe how to translate a logical 3‑D index into a linear offset from that base.
- tensor_of_strides_abc = cute.make_tensor(
- ptr_of_tensor_strides_abc, cute.make_layout((num_groups, 3, 2), stride=(6, 2, 1))
- )
-
# Use fake shape for initial Tma descriptor and atom setup
# The real Tma desc and atom will be updated during kernel execution.
min_a_shape = (cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(64), cutlass.Int32(1))
⋯ 191 unchanged lines
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]
+ problem_sizes_buffer: cute.struct.MemRange[cutlass.Int32, num_groups * 4]
+ ptrs_abc_buffer: cute.struct.MemRange[cutlass.Int64, num_groups * 3]
+ ptrs_sfasfb_buffer: cute.struct.MemRange[cutlass.Int64, num_groups * 2]
+ group_tile_prefix_end: cute.struct.MemRange[cutlass.Int32, num_groups]
+ group_cluster_count_m: cute.struct.MemRange[cutlass.Int32, num_groups]
+ group_cluster_count_k: cute.struct.MemRange[cutlass.Int32, num_groups]
# tmem_dealloc_mbar_ptr is the shared‑memory barrier (mbarrier) used to synchronize TMEM deallocation across warps/CTAs.
# It ensures everyone is done using tensor memory before it’s freed.
⋯ 182 unchanged lines
ptrs_abc=tensor_of_abc_ptrs,
ptrs_sfasfb=tensor_of_sfasfb_ptrs,
tensormaps=tensor_of_tensormap,
- strides_abc=tensor_of_strides_abc
- # might be just this So if you choose A=MKL, B=NKL, C=MNL (i.e., A is M‑major, B is N‑major, C is M‑major), then yes—is_mode0_major is True for all three, and the stride for each group becomes (1, mode0).
- # strides_abc,
- # tensor_address_abc,
- # tensor_address_sfasfb,
- # tensormap_cute_tensor,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
⋯ 9 unchanged lines
group_idx: cutlass.Int32,
dtype: Type[cutlass.Numeric],
problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32],
- strides_abc: cute.Tensor,
tensor_address_abc: cute.Tensor,
tensor_index: int,
):
- """Extract stride and tensor address for a given group and construct a global tensor for A, B or C.
+ """Construct a global tensor view for A, B or C for a specific group.
This function is used within the kernel to dynamically create a CUTE tensor
representing A, B, or C for the current group being processed, using the
- group-specific address, shape, and stride information.
+ group-specific address and shape information.
:param group_idx: The index of the current group within the grouped GEMM.
:type group_idx: cutlass.Int32
⋯ 1 unchanged lines
:type dtype: Type[cutlass.Numeric]
:param problem_shape_mnk: The (M, N, K) problem shape for the current group.
:type problem_shape_mnk: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32]
- :param strides_abc: Tensor containing strides for A, B, C for all groups. Layout: (group_count, 3, 2).
- :type strides_abc: cute.Tensor
:param tensor_address_abc: Tensor containing global memory addresses for A, B, C for all groups. Layout: (group_count, 3).
:type tensor_address_abc: cute.Tensor
:param tensor_index: Specifies which tensor to create: 0 for A, 1 for B, 2 for C.
⋯ 13 unchanged lines
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 = strides_tensor_reg[0]
- stride_k = strides_tensor_reg[1]
c1 = cutlass.Int32(1)
c0 = cutlass.Int32(0)
⋯ 2 unchanged lines
k = problem_shape_mnk[2]
return cute.make_tensor(
tensor_gmem_ptr,
- cute.make_layout((m, k, c1), stride=(stride_mn, stride_k, c0)),
+ cute.make_layout((m, k, c1), stride=(k, c1, c0)),
)
elif cutlass.const_expr(tensor_index == 1): # tensor B
n = problem_shape_mnk[1]
k = problem_shape_mnk[2]
return cute.make_tensor(
tensor_gmem_ptr,
- cute.make_layout((n, k, c1), stride=(stride_mn, stride_k, c0)),
+ cute.make_layout((n, k, c1), stride=(k, c1, c0)),
)
else: # tensor C
m = problem_shape_mnk[0]
n = problem_shape_mnk[1]
return cute.make_tensor(
tensor_gmem_ptr,
- cute.make_layout((m, n, c1), stride=(stride_mn, stride_k, c0)),
+ cute.make_layout((m, n, c1), stride=(n, c1, c0)),
)
@cute.jit
⋯ 167 unchanged lines
- occupancy * (mbar_helpers_bytes + c_bytes)
) // (occupancy * c_bytes_per_stage)
+ # Optional manual tuning knobs.
+ if max_num_ab_stage > 0:
+ num_ab_stage = min(num_ab_stage, max_num_ab_stage)
+ if force_num_c_stage > 0:
+ num_c_stage = force_num_c_stage
+
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
⋯ 26 unchanged lines
# (128, 128, 352)
tile_sched_params = utils.PersistentTileSchedulerParams(
- problem_shape_ntile_mnl, (*cluster_shape_mn, 1)
+ problem_shape_ntile_mnl,
+ (*cluster_shape_mn, 1),
+ swizzle_size=tile_swizzle_size,
+ raster_along_m=tile_raster_along_m,
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
⋯ 225 unchanged lines
_tensormap_workspace_cache = {}
_pointer_metadata_cache = {}
_launch_context_cache = {}
+ _launch_context_fast_cache = {}
_hardware_info = None
_sm_count_cache = None
_max_active_clusters_cache = {}
⋯ 32 unchanged lines
def _get_static_metadata_tensors(
problem_sizes: List[Tuple[int, int, int, int]],
- ) -> Tuple[torch.Tensor, torch.Tensor]:
+ ) -> torch.Tensor:
device_idx = torch.cuda.current_device()
cache_key = (device_idx, make_problem_sizes_cache_key(problem_sizes))
cached = _static_metadata_cache.get(cache_key)
⋯ 4 unchanged lines
tensor_of_problem_sizes = torch.tensor(
problem_sizes, dtype=torch.int32, device="cuda"
)
- # (num_groups, 3, 2)
- strides_abc = [[(k, 1), (k, 1), (n, 1)] for _, n, k, _ in problem_sizes]
- tensor_of_strides_abc = torch.tensor(strides_abc, dtype=torch.int32, device="cuda")
+ _static_metadata_cache[cache_key] = tensor_of_problem_sizes
+ return tensor_of_problem_sizes
- _static_metadata_cache[cache_key] = (
- tensor_of_problem_sizes,
- tensor_of_strides_abc,
- )
- return tensor_of_problem_sizes, tensor_of_strides_abc
-
def _get_pointer_metadata_tensors(
abc_ptrs,
sfasfb_ptrs,
⋯ 68 unchanged lines
cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16,
)
- cute_ptr_of_tensor_of_strides_abc = make_ptr(
- cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16,
- )
-
kernel_runner = Sm100GroupedBlockScaledGemmKernel(
- (128, 128),
+ (mma_tiler_mnk[0], mma_tiler_mnk[1]),
cluster_shape_mn=cluster_shape_mn
)
⋯ 1 unchanged lines
# Fake cluster numbers for compile only.
# Compute the tile shape for each CUDA Thread Block (CTA)
# cta_tile_shape_mn: [M_tile, N_tile] = [128, 128] for this kernel
- cta_tile_shape_mn = [128, mma_tiler_mnk[1]]
+ cta_tile_shape_mn = [mma_tiler_mnk[0], mma_tiler_mnk[1]]
# cluster_tile_shape_mn: Total tile shape per cluster
cluster_tile_shape_mn = tuple(
x * y for x, y in zip(cta_tile_shape_mn, cluster_shape_mn)
⋯ 21 unchanged lines
cute_ptr_of_tensor_of_abc_ptrs,
cute_ptr_of_tensor_of_sfasfb_ptrs,
cute_ptr_of_tensor_of_tensormap,
- cute_ptr_of_tensor_of_strides_abc,
sm_count,
total_num_clusters,
num_groups,
⋯ 31 unchanged lines
list of c tensors where c is torch.Tensor[float16] of shape [m, n, l] for each group
"""
abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
+ device_idx = torch.cuda.current_device()
+ problem_key = make_problem_sizes_cache_key(problem_sizes)
- abc_ptrs = tuple((a.data_ptr(), b.data_ptr(), c.data_ptr()) for a, b, c in abc_tensors)
- sfasfb_ptrs = tuple(
- (sfa_reordered.data_ptr(), sfb_reordered.data_ptr())
- for sfa_reordered, sfb_reordered in sfasfb_reordered_tensors
+ # Fast path cache keyed by object identity. Guard against Python id reuse by
+ # verifying referential identity before trusting the cached launch context.
+ fast_cache_key = (
+ device_idx,
+ problem_key,
+ id(abc_tensors),
+ id(sfasfb_reordered_tensors),
)
- launch_cache_key = (
- torch.cuda.current_device(),
- make_problem_sizes_cache_key(problem_sizes),
- abc_ptrs,
- sfasfb_ptrs,
- )
- cached_launch = _launch_context_cache.get(launch_cache_key)
+ cached_fast = _launch_context_fast_cache.get(fast_cache_key)
+ if cached_fast is not None:
+ cached_abc_obj, cached_sfa_obj, cached_launch = cached_fast
+ if cached_abc_obj is not abc_tensors or cached_sfa_obj is not sfasfb_reordered_tensors:
+ _launch_context_fast_cache.pop(fast_cache_key, None)
+ cached_launch = None
+ else:
+ cached_launch = None
+
if cached_launch is None:
+ abc_ptrs = tuple((a.data_ptr(), b.data_ptr(), c.data_ptr()) for a, b, c in abc_tensors)
+ sfasfb_ptrs = tuple(
+ (sfa_reordered.data_ptr(), sfb_reordered.data_ptr())
+ for sfa_reordered, sfb_reordered in sfasfb_reordered_tensors
+ )
+ launch_cache_key = (
+ device_idx,
+ problem_key,
+ abc_ptrs,
+ sfasfb_ptrs,
+ )
+ cached_launch = _launch_context_cache.get(launch_cache_key)
+ if cached_launch is None:
compiled_func = compile_kernel(problem_sizes)
- tensor_of_problem_sizes, tensor_of_strides_abc = _get_static_metadata_tensors(
- problem_sizes
- )
+ tensor_of_problem_sizes = _get_static_metadata_tensors(problem_sizes)
# Cache pointer tensors to avoid per-call CPU work and H2D metadata upload.
# tensor_of_abc_ptrs: Shape (num_groups, 3) containing (a_ptr, b_ptr, c_ptr) per group
⋯ 37 unchanged lines
cute.AddressSpace.gmem,
assumed_align=16,
)
- cute_ptr_of_tensor_of_strides_abc = make_ptr(
- cutlass.Int32,
- tensor_of_strides_abc.data_ptr(),
- cute.AddressSpace.gmem,
- assumed_align=16,
- )
-
cached_launch = (
compiled_func,
cute_ptr_of_tensor_of_problem_sizes,
cute_ptr_of_tensor_of_abc_ptrs,
cute_ptr_of_tensor_of_sfasfb_ptrs,
cute_ptr_of_tensor_of_tensormap,
- cute_ptr_of_tensor_of_strides_abc,
)
_launch_context_cache[launch_cache_key] = cached_launch
- else:
- (
- compiled_func,
- cute_ptr_of_tensor_of_problem_sizes,
- cute_ptr_of_tensor_of_abc_ptrs,
- cute_ptr_of_tensor_of_sfasfb_ptrs,
- cute_ptr_of_tensor_of_tensormap,
- cute_ptr_of_tensor_of_strides_abc,
- ) = cached_launch
+ _launch_context_fast_cache[fast_cache_key] = (
+ abc_tensors,
+ sfasfb_reordered_tensors,
+ cached_launch,
+ )
+ (
+ compiled_func,
+ cute_ptr_of_tensor_of_problem_sizes,
+ cute_ptr_of_tensor_of_abc_ptrs,
+ cute_ptr_of_tensor_of_sfasfb_ptrs,
+ cute_ptr_of_tensor_of_tensormap,
+ ) = cached_launch
# Launch the JIT-compiled GPU kernel with all prepared data
# The kernel will perform block-scaled group GEMM: C = A * SFA * B * SFB for all groups
⋯ 2 unchanged lines
cute_ptr_of_tensor_of_abc_ptrs, # Pointer to ABC tensor pointers array
cute_ptr_of_tensor_of_sfasfb_ptrs, # Pointer to scale factor pointers array
cute_ptr_of_tensor_of_tensormap, # Pointer to tensormap buffer
- cute_ptr_of_tensor_of_strides_abc, # Pointer to tensor
)
num_groups = len(problem_sizes)
scrolls · 764 diff lines total

Best evidence level for this revision: reported

JSON