Skip to content
KernelIndex
Search⌘K

submission 488633

yolojewjitsu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-488633?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
32.9µs
#176 of 310
2026-02-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5ab8798f35d3d486c28c7042584ddcc4709daf1432ddfa7db2e8baec74f570ae
license declaredunknown
license concludedunknown
authorsyolojewjitsu
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(
mbarrierepilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(epilog_warp_ids))
persistent-kerneltile_sched_params: utils.PersistentTileSchedulerParams,
shared-memorysmem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
tcgen05cta_group = tcgen05.CtaGroup.ONE
warp-specializationnum_tma_producer = num_mcast_ctas_a + num_mcast_ctas_b - 1

Kernel source

mission.py890 lines
import functools
from typing import List, Tuple, Union
from inspect import isclass

import torch

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

from task import input_t, output_t

ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
acc_dtype = cutlass.Float32
sf_vec_size = 16
mma_tiler_mn = (128, 128)
cluster_shape_mn = (1, 2)

mma_inst_shape_k = 64
mma_inst_tile_k = 4
mma_tiler_mnk = (*mma_tiler_mn, mma_inst_shape_k * mma_inst_tile_k)

bytes_per_tensormap = 128
num_tensormaps = 5  
epilog_warp_ids = (0, 1, 2, 3)
mma_warp_id = 4
tma_warp_id = 5
threads_per_cta = 32 * 6 


epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(epilog_warp_ids))
tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=2, num_threads=32 * (1 + len(epilog_warp_ids)))
tensormap_ab_init_barrier = pipeline.NamedBarrier(barrier_id=3, num_threads=64)

SM100_TMEM_CAPACITY_COLUMNS = 512
smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")

def compute_stages(c_layout):
    """Compute num_acc_stage, num_ab_stage, num_c_stage."""
    cta_group = tcgen05.CtaGroup.ONE
    tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
        ab_dtype, "M", "N", sf_dtype, sf_vec_size, cta_group, mma_tiler_mn,
    )
    num_acc_stage = 2

    epi_tile = sm100_utils.compute_epilogue_tile_shape(
        (mma_tiler_mnk[0], mma_tiler_mnk[1], mma_tiler_mnk[2]),
        False,
        c_layout,
        c_dtype,
    )
    num_c_stage = 2

    a_smem_1 = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, ab_dtype, 1)
    b_smem_1 = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, ab_dtype, 1)
    sfa_smem_1 = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
    sfb_smem_1 = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
    c_smem_1 = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)

    ab_bytes = (
        cute.size_in_bytes(ab_dtype, a_smem_1)
        + cute.size_in_bytes(ab_dtype, b_smem_1)
        + cute.size_in_bytes(sf_dtype, sfa_smem_1)
        + cute.size_in_bytes(sf_dtype, sfb_smem_1)
    )
    c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_1)
    reserved = 1024
    occupancy = 1
    c_bytes = c_bytes_per_stage * num_c_stage

    num_ab_stage = (smem_capacity // occupancy - (reserved + c_bytes)) // ab_bytes

    num_c_stage += (
        smem_capacity - occupancy * ab_bytes * num_ab_stage
        - occupancy * (reserved + c_bytes)
    ) // (occupancy * c_bytes_per_stage)

    return num_acc_stage, num_ab_stage, num_c_stage, epi_tile

@cute.jit
def make_real_tensor_ab(ptr_i64, dtype, outer_dim, k):
    """Construct global tensor for A (m,k,1) or B (n,k,1) from pointer and dimensions."""
    tensor_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
    c1 = cutlass.Int32(1)
    c0 = cutlass.Int32(0)
    return cute.make_tensor(
        tensor_ptr,
        cute.make_layout(
            (outer_dim, k, c1),
            stride=(k, cutlass.Int32(1), c0),
        ),
    )


@cute.jit
def make_real_tensor_c(ptr_i64, m, n):
    """Construct global tensor for C from pointer and dimensions."""
    tensor_ptr = cute.make_ptr(c_dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
    c1 = cutlass.Int32(1)
    c0 = cutlass.Int32(0)
    return cute.make_tensor(
        tensor_ptr,
        cute.make_layout(
            (m, n, c1),
            stride=(n, cutlass.Int32(1), c0),
        ),
    )


@cute.jit
def make_real_tensor_sf(ptr_i64, outer_dim, k):
    """Construct global tensor for SFA or SFB from pointer and dimensions."""
    tensor_ptr = cute.make_ptr(sf_dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
    c1 = cutlass.Int32(1)
    sf_layout = blockscaled_utils.tile_atom_to_shape_SF(
        (outer_dim, k, c1), sf_vec_size
    )
    return cute.make_tensor(tensor_ptr, sf_layout)

@cute.kernel
def kernel(
    tiled_mma: 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,
    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,
    num_tma_load_bytes: cutlass.Constexpr[int],
    num_ab_stage: cutlass.Constexpr[int],
    num_acc_stage: cutlass.Constexpr[int],
    num_c_stage: cutlass.Constexpr[int],
    cluster_tile_shape_mnk: cutlass.Constexpr,
    c_layout: cutlass.Constexpr,
):
    warp_idx = cute.arch.warp_idx()
    warp_idx = cute.arch.make_warp_uniform(warp_idx)

    if warp_idx == tma_warp_id:
        cpasync.prefetch_descriptor(tma_atom_a)
        cpasync.prefetch_descriptor(tma_atom_b)
        cpasync.prefetch_descriptor(tma_atom_sfa)
        cpasync.prefetch_descriptor(tma_atom_sfb)
        cpasync.prefetch_descriptor(tma_atom_c)

    bidx, bidy, bidz = cute.arch.block_idx()
    tidx, _, _ = cute.arch.thread_idx()

    size_tensormap_in_i64 = num_tensormaps * bytes_per_tensormap // 8

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

    smem = utils.SmemAllocator()
    storage = smem.allocate(SharedStorage)

    tensormap_smem_ptr = storage.tensormap_buffer.data_ptr()
    tm_step = bytes_per_tensormap // 8
    tensormap_a_smem_ptr = tensormap_smem_ptr
    tensormap_b_smem_ptr = tensormap_a_smem_ptr + tm_step
    tensormap_sfa_smem_ptr = tensormap_b_smem_ptr + tm_step
    tensormap_sfb_smem_ptr = tensormap_sfa_smem_ptr + tm_step
    tensormap_c_smem_ptr = tensormap_sfb_smem_ptr + tm_step

    tmem_holding_buf = storage.tmem_holding_buf

    cluster_layout_vmnk = cute.tiled_divide(cute.make_layout((*cluster_shape_mn, 1)), (tiled_mma.thr_id.shape,))

    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)

    num_mcast_ctas_a = cute.size(cluster_layout_vmnk.shape[2])
    num_mcast_ctas_b = cute.size(cluster_layout_vmnk.shape[1])
    num_tma_producer = num_mcast_ctas_a + num_mcast_ctas_b - 1

    ab_pipeline = pipeline.PipelineTmaUmma.create(
        barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
        num_stages=num_ab_stage,
        producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
        consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, num_tma_producer),
        tx_count=num_tma_load_bytes,
        cta_layout_vmnk=cluster_layout_vmnk,
    )
    acc_pipeline = pipeline.PipelineUmmaAsync.create(
        barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
        num_stages=num_acc_stage,
        producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
        consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, len(epilog_warp_ids)),
        cta_layout_vmnk=cluster_layout_vmnk,
    )

    pipeline_init_arrive(cluster_shape_mn=cluster_shape_mn, is_relaxed=True)

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

    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_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1)

    gA_mkl = cute.local_tile(mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None))
    gB_nkl = cute.local_tile(mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None))
    gSFA_mkl = cute.local_tile(mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None))
    gSFB_nkl = cute.local_tile(mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None))
    gC_mnl = cute.local_tile(mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None))

    mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
    thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
    tCgA = thr_mma.partition_A(gA_mkl)
    tCgB = thr_mma.partition_B(gB_nkl)
    tCgSFA = thr_mma.partition_A(gSFA_mkl)
    tCgSFB = thr_mma.partition_B(gSFB_nkl)
    tCgC = thr_mma.partition_C(gC_mnl)

    a_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape)
    b_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape)

    tAsA, tAgA = cpasync.tma_partition(tma_atom_a, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sA, 0, 3), cute.group_modes(tCgA, 0, 3))
    tBsB, tBgB = cpasync.tma_partition(tma_atom_b, block_in_cluster_coord_vmnk[1], b_cta_layout, cute.group_modes(sB, 0, 3), cute.group_modes(tCgB, 0, 3))
    tAsSFA, tAgSFA = cpasync.tma_partition(tma_atom_sfa, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sSFA, 0, 3), cute.group_modes(tCgSFA, 0, 3))
    tAsSFA = cute.filter_zeros(tAsSFA)
    tAgSFA = cute.filter_zeros(tAgSFA)
    tBsSFB, tBgSFB = cpasync.tma_partition(tma_atom_sfb, block_in_cluster_coord_vmnk[1], b_cta_layout, cute.group_modes(sSFB, 0, 3), cute.group_modes(tCgSFB, 0, 3))
    tBsSFB = cute.filter_zeros(tBsSFB)
    tBgSFB = cute.filter_zeros(tBgSFB)

    tCrA = tiled_mma.make_fragment_A(sA)
    tCrB = tiled_mma.make_fragment_B(sB)
    acc_shape = tiled_mma.partition_shape_C(mma_tiler_mn)
    tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, num_acc_stage))

    pipeline_init_wait(cluster_shape_mn=cluster_shape_mn)

    grid_dim = cute.arch.grid_dim()
    tensormap_workspace_idx = bidz * grid_dim[1] * grid_dim[0] + bidy * grid_dim[0] + bidx

    tensormap_manager = utils.TensorMapManager(utils.TensorMapUpdateMode.SMEM, 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)

    if warp_idx == tma_warp_id:
        tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
        group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
            group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
        )
        tensormap_init_done = cutlass.Boolean(False)
        last_group_idx = cutlass.Int32(-1)
        work_tile = tile_sched.initial_work_tile_info()
        ab_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, num_ab_stage)

        while work_tile.is_valid_tile:
            cur_tile_coord = work_tile.tile_idx
            grouped_gemm_info = group_gemm_ts_helper.delinearize_z(cur_tile_coord, problem_sizes_mnkl)
            cur_k_tile_cnt = grouped_gemm_info.cta_tile_count_k
            cur_group_idx = grouped_gemm_info.group_idx
            is_group_changed = cur_group_idx != last_group_idx

            if is_group_changed:
                m_cur = grouped_gemm_info.problem_shape_m
                n_cur = grouped_gemm_info.problem_shape_n
                k_cur = grouped_gemm_info.problem_shape_k

                real_a = make_real_tensor_ab(ptrs_abc[(cur_group_idx, 0)], ab_dtype, m_cur, k_cur)
                real_b = make_real_tensor_ab(ptrs_abc[(cur_group_idx, 1)], ab_dtype, n_cur, k_cur)
                real_sfa = make_real_tensor_sf(ptrs_sfasfb[(cur_group_idx, 0)], m_cur, k_cur)
                real_sfb = make_real_tensor_sf(ptrs_sfasfb[(cur_group_idx, 1)], n_cur, k_cur)

                if tensormap_init_done == False:
                    tensormap_ab_init_barrier.arrive_and_wait()
                    tensormap_init_done = True

                tensormap_manager.update_tensormap(
                    (real_a, real_b, real_sfa, real_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),
                    tma_warp_id,
                    (tensormap_a_smem_ptr, tensormap_b_smem_ptr, tensormap_sfa_smem_ptr, tensormap_sfb_smem_ptr),
                )

            mma_tile_coord_mnl = (
                grouped_gemm_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
                grouped_gemm_info.cta_tile_idx_n,
                0,
            )

            tAgA_s = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
            tBgB_s = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
            tAgSFA_s = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
            tBgSFB_s = tBgSFB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]

            ab_producer_state.reset_count()
            peek_status = cutlass.Boolean(1)
            if ab_producer_state.count < cur_k_tile_cnt:
                peek_status = ab_pipeline.producer_try_acquire(ab_producer_state)

            if is_group_changed:
                tensormap_manager.fence_tensormap_update(tensormap_a_gmem_ptr)
                tensormap_manager.fence_tensormap_update(tensormap_b_gmem_ptr)
                tensormap_manager.fence_tensormap_update(tensormap_sfa_gmem_ptr)
                tensormap_manager.fence_tensormap_update(tensormap_sfb_gmem_ptr)

            for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
                ab_pipeline.producer_acquire(ab_producer_state, peek_status)

                cute.copy(tma_atom_a, tAgA_s[(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=tensormap_manager.get_tensormap_ptr(tensormap_a_gmem_ptr, cute.AddressSpace.generic))
                cute.copy(tma_atom_b, tBgB_s[(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=tensormap_manager.get_tensormap_ptr(tensormap_b_gmem_ptr, cute.AddressSpace.generic))
                cute.copy(tma_atom_sfa, tAgSFA_s[(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=tensormap_manager.get_tensormap_ptr(tensormap_sfa_gmem_ptr, cute.AddressSpace.generic))
                cute.copy(tma_atom_sfb, tBgSFB_s[(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=tensormap_manager.get_tensormap_ptr(tensormap_sfb_gmem_ptr, cute.AddressSpace.generic))

                ab_producer_state.advance()
                peek_status = cutlass.Boolean(1)
                if ab_producer_state.count < cur_k_tile_cnt:
                    peek_status = ab_pipeline.producer_try_acquire(ab_producer_state)

            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()
            last_group_idx = cur_group_idx

        ab_pipeline.producer_tail(ab_producer_state)

    if warp_idx == mma_warp_id:
        tensormap_manager.init_tensormap_from_atom(tma_atom_a, tensormap_a_smem_ptr, mma_warp_id)
        tensormap_manager.init_tensormap_from_atom(tma_atom_b, tensormap_b_smem_ptr, mma_warp_id)
        tensormap_manager.init_tensormap_from_atom(tma_atom_sfa, tensormap_sfa_smem_ptr, mma_warp_id)
        tensormap_manager.init_tensormap_from_atom(tma_atom_sfb, tensormap_sfb_smem_ptr, mma_warp_id)
        tensormap_ab_init_barrier.arrive_and_wait()

        tmem_alloc_barrier.arrive_and_wait()

        acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
        tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

        sfa_tmem_ptr = cute.recast_ptr(acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base), dtype=sf_dtype)
        tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)))
        tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)

        sfb_tmem_ptr = cute.recast_ptr(acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base) + tcgen05.find_tmem_tensor_col_offset(tCtSFA), dtype=sf_dtype)
        tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)))
        tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)

        copy_atom_s2t = cute.make_copy_atom(tcgen05.Cp4x32x128bOp(tcgen05.CtaGroup.ONE), sf_dtype)

        tCsSFA_compact = cute.filter_zeros(sSFA)
        tCtSFA_compact = cute.filter_zeros(tCtSFA)
        tiled_copy_s2t_sfa = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFA_compact)
        thr_s2t_sfa = tiled_copy_s2t_sfa.get_slice(0)
        tCsSFA_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t_sfa, thr_s2t_sfa.partition_S(tCsSFA_compact))
        tCtSFA_s2t = thr_s2t_sfa.partition_D(tCtSFA_compact)

        tCsSFB_compact = cute.filter_zeros(sSFB)
        tCtSFB_compact = cute.filter_zeros(tCtSFB)
        tiled_copy_s2t_sfb = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFB_compact)
        thr_s2t_sfb = tiled_copy_s2t_sfb.get_slice(0)
        tCsSFB_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t_sfb, thr_s2t_sfb.partition_S(tCsSFB_compact))
        tCtSFB_s2t = thr_s2t_sfb.partition_D(tCtSFB_compact)

        tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
        group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
            group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
        )
        work_tile = tile_sched.initial_work_tile_info()
        ab_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, num_ab_stage)
        acc_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, num_acc_stage)

        while work_tile.is_valid_tile:
            cur_tile_coord = work_tile.tile_idx
            cur_k_tile_cnt, cur_group_idx = group_gemm_ts_helper.search_cluster_tile_count_k(cur_tile_coord, problem_sizes_mnkl)

            tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]

            ab_consumer_state.reset_count()
            peek_status = cutlass.Boolean(1)
            if ab_consumer_state.count < cur_k_tile_cnt:
                peek_status = ab_pipeline.consumer_try_wait(ab_consumer_state)

            acc_pipeline.producer_acquire(acc_producer_state)
            tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

            for k_tile in range(cur_k_tile_cnt):
                ab_pipeline.consumer_wait(ab_consumer_state, peek_status)

                s2t_coord = (None, None, None, None, ab_consumer_state.index)
                cute.copy(tiled_copy_s2t_sfa, tCsSFA_s2t[s2t_coord], tCtSFA_s2t)
                cute.copy(tiled_copy_s2t_sfb, tCsSFB_s2t[s2t_coord], tCtSFB_s2t)

                num_kblocks = cute.size(tCrA, mode=[2])
                for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                    kblock_coord = (None, None, kblock_idx, ab_consumer_state.index)
                    sf_coord = (None, None, kblock_idx)
                    tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_coord].iterator)
                    tiled_mma.set(tcgen05.Field.SFB, tCtSFB[sf_coord].iterator)
                    cute.gemm(tiled_mma, tCtAcc, tCrA[kblock_coord], tCrB[kblock_coord], tCtAcc)
                    tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

                ab_pipeline.consumer_release(ab_consumer_state)
                ab_consumer_state.advance()
                peek_status = cutlass.Boolean(1)
                if ab_consumer_state.count < cur_k_tile_cnt:
                    peek_status = ab_pipeline.consumer_try_wait(ab_consumer_state)

            acc_pipeline.producer_commit(acc_producer_state)
            acc_producer_state.advance()

            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()

        acc_pipeline.producer_tail(acc_producer_state)

    if warp_idx < mma_warp_id:
        tensormap_manager.init_tensormap_from_atom(tma_atom_c, tensormap_c_smem_ptr, epilog_warp_ids[0])

        if warp_idx == epilog_warp_ids[0]:
            cute.arch.alloc_tmem(SM100_TMEM_CAPACITY_COLUMNS, tmem_holding_buf, is_two_cta=False)

        tmem_alloc_barrier.arrive_and_wait()

        acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
        tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

        epi_tidx = tidx
        copy_atom_t2r = sm100_utils.get_tmem_load_op(
            (mma_tiler_mnk[0], mma_tiler_mnk[1], mma_tiler_mnk[2]),
            c_layout, c_dtype, acc_dtype, epi_tile, False,
        )
        tAcc_epi = cute.flat_divide(tCtAcc_base[((None, None), 0, 0, None)], epi_tile)
        tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)])
        thr_copy_t2r = tiled_copy_t2r.get_slice(epi_tidx)
        tTR_tAcc_base = thr_copy_t2r.partition_S(tAcc_epi)

        gC_mnl_epi = cute.flat_divide(tCgC[((None, None), 0, 0, None, None, None)], epi_tile)
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        tTR_rAcc = cute.make_rmem_tensor(tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, acc_dtype)

        tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, c_dtype)
        copy_atom_r2s = sm100_utils.get_smem_store_op(c_layout, c_dtype, acc_dtype, tiled_copy_t2r)
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
        thr_copy_r2s = tiled_copy_r2s.get_slice(epi_tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        tRS_rC = tiled_copy_r2s.retile(tTR_rC)

        sC_for_tma = cute.group_modes(sC, 0, 2)
        gC_for_tma = cute.group_modes(gC_mnl_epi, 0, 2)
        bSG_sC, bSG_gC_partitioned = cpasync.tma_partition(tma_atom_c, 0, cute.make_layout(1), sC_for_tma, gC_for_tma)

        tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
        group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
            group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
        )
        work_tile = tile_sched.initial_work_tile_info()
        acc_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, num_acc_stage)

        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=num_c_stage,
            producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 32 * len(epilog_warp_ids)),
        )
        last_group_idx = cutlass.Int32(-1)

        while work_tile.is_valid_tile:
            cur_tile_coord = work_tile.tile_idx
            grouped_gemm_info = group_gemm_ts_helper.delinearize_z(cur_tile_coord, problem_sizes_mnkl)
            cur_group_idx = grouped_gemm_info.group_idx
            is_group_changed = cur_group_idx != last_group_idx

            if is_group_changed:
                m_cur = grouped_gemm_info.problem_shape_m
                n_cur = grouped_gemm_info.problem_shape_n
                k_cur = grouped_gemm_info.problem_shape_k
                real_c = make_real_tensor_c(ptrs_abc[(cur_group_idx, 2)], m_cur, n_cur)
                tensormap_manager.update_tensormap(
                    ((real_c),), ((tma_atom_c),), ((tensormap_c_gmem_ptr),),
                    epilog_warp_ids[0], (tensormap_c_smem_ptr,),
                )

            mma_tile_coord_mnl = (
                grouped_gemm_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
                grouped_gemm_info.cta_tile_idx_n,
                0,
            )

            bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]

            tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_consumer_state.index)]
            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 == epilog_warp_ids[0]:
                    tensormap_manager.fence_tensormap_update(tensormap_c_gmem_ptr)

            subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
            num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
            for subtile_idx in range(subtile_cnt):
                cute.copy(tiled_copy_t2r, tTR_tAcc[(None, None, None, subtile_idx)], tTR_rAcc)

                acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
                tRS_rC.store(acc_vec.to(c_dtype))

                c_buffer = (num_prev_subtiles + subtile_idx) % num_c_stage
                cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
                cute.arch.fence_proxy("async.shared", space="cta")
                epilog_sync_barrier.arrive_and_wait()

                if warp_idx == epilog_warp_ids[0]:
                    cute.copy(
                        tma_atom_c,
                        bSG_sC[(None, c_buffer)],
                        bSG_gC[(None, subtile_idx)],
                        tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_c_gmem_ptr, cute.AddressSpace.generic),
                    )
                    c_pipeline.producer_commit()
                    c_pipeline.producer_acquire()
                epilog_sync_barrier.arrive_and_wait()

            with cute.arch.elect_one():
                acc_pipeline.consumer_release(acc_consumer_state)
            acc_consumer_state.advance()

            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()
            last_group_idx = cur_group_idx

        if warp_idx == epilog_warp_ids[0]:
            cute.arch.relinquish_tmem_alloc_permit(is_two_cta=False)
        epilog_sync_barrier.arrive_and_wait()
        if warp_idx == epilog_warp_ids[0]:
            cute.arch.dealloc_tmem(acc_tmem_ptr, SM100_TMEM_CAPACITY_COLUMNS, is_two_cta=False)
        c_pipeline.producer_tail()

@cute.jit
def launch_kernel(
    ptr_problem_sizes: cute.Pointer,
    ptr_abc: cute.Pointer,
    ptr_sfasfb: cute.Pointer,
    ptr_tensormap: cute.Pointer,
    total_num_clusters: cutlass.Constexpr[int],
    num_tensormap_entries: cutlass.Constexpr[int],
    problem_sizes: List[Tuple[int, int, int, int]],
    max_active_clusters: cutlass.Constexpr[int],
):
    min_shape = (cutlass.Int32(128), cutlass.Int32(128), cutlass.Int32(256), cutlass.Int32(1))
    null_ab = cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    null_sf = cute.make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    null_c = cute.make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)

    initial_a = cute.make_tensor(null_ab, cute.make_layout(
        (min_shape[0], cute.assume(min_shape[2], 32), min_shape[3]),
        stride=(cute.assume(min_shape[2], 32), 1, cute.assume(min_shape[0] * min_shape[2], 32)),
    ))
    initial_b = cute.make_tensor(null_ab, cute.make_layout(
        (min_shape[1], cute.assume(min_shape[2], 32), min_shape[3]),
        stride=(cute.assume(min_shape[2], 32), 1, cute.assume(min_shape[1] * min_shape[2], 32)),
    ))
    initial_c = cute.make_tensor(null_c, cute.make_layout(
        (cute.assume(min_shape[0], 32), min_shape[1], min_shape[3]),
        stride=(min_shape[1], 1, cute.assume(min_shape[0] * min_shape[1], 32)),
    ))

    c_layout = utils.LayoutEnum.from_tensor(initial_c)

    NUM_ACC_STAGE, NUM_AB_STAGE, NUM_C_STAGE, _ = compute_stages(c_layout)

    num_groups = len(problem_sizes)
    tensor_problem_sizes = cute.make_tensor(ptr_problem_sizes, cute.make_layout((num_groups, 4), stride=(4, 1)))
    tensor_abc = cute.make_tensor(ptr_abc, cute.make_layout((num_groups, 3), stride=(3, 1)))
    tensor_sfasfb = cute.make_tensor(ptr_sfasfb, cute.make_layout((num_groups, 2), stride=(2, 1)))
    tensor_tensormap = cute.make_tensor(ptr_tensormap, cute.make_layout((num_tensormap_entries, num_tensormaps, bytes_per_tensormap // 8), stride=(num_tensormaps * bytes_per_tensormap // 8, bytes_per_tensormap // 8, 1)))

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

    cta_group = tcgen05.CtaGroup.ONE
    tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
        ab_dtype, "M", "N", sf_dtype, sf_vec_size, cta_group, mma_tiler_mn,
    )
    atom_thr_size = cute.size(tiled_mma.thr_id.shape)

    cluster_layout_vmnk = cute.tiled_divide(cute.make_layout((*cluster_shape_mn, 1)), (tiled_mma.thr_id.shape,))

    is_a_mcast = cluster_shape_mn[1] > 1
    is_b_mcast = cluster_shape_mn[0] > 1
    tma_op_a = cpasync.CopyBulkTensorTileG2SMulticastOp(cta_group) if is_a_mcast else cpasync.CopyBulkTensorTileG2SOp(cta_group)
    tma_op_b = cpasync.CopyBulkTensorTileG2SMulticastOp(cta_group) if is_b_mcast else cpasync.CopyBulkTensorTileG2SOp(cta_group)

    a_smem_staged = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, ab_dtype, NUM_AB_STAGE)
    b_smem_staged = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, ab_dtype, NUM_AB_STAGE)
    sfa_smem_staged = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, NUM_AB_STAGE)
    sfb_smem_staged = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, NUM_AB_STAGE)

    epi_tile = sm100_utils.compute_epilogue_tile_shape(
        mma_tiler_mnk, False, c_layout, c_dtype,
    )
    c_smem_staged = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, NUM_C_STAGE)

    a_smem = cute.slice_(a_smem_staged, (None, None, None, 0))
    tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
        tma_op_a, initial_a, a_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
    )
    b_smem = cute.slice_(b_smem_staged, (None, None, None, 0))
    tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
        tma_op_b, initial_b, b_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
    )
    sfa_smem = cute.slice_(sfa_smem_staged, (None, None, None, 0))
    tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
        tma_op_a, initial_sfa, sfa_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )
    sfb_smem = cute.slice_(sfb_smem_staged, (None, None, None, 0))
    tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
        tma_op_b, initial_sfb, sfb_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )

    epi_smem = cute.slice_(c_smem_staged, (None, None, 0))
    tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
        cpasync.CopyBulkTensorTileS2GOp(), initial_c, epi_smem, epi_tile,
    )

    num_tma_load_bytes = (
        cute.size_in_bytes(ab_dtype, a_smem)
        + cute.size_in_bytes(ab_dtype, b_smem)
        + cute.size_in_bytes(sf_dtype, sfa_smem)
        + cute.size_in_bytes(sf_dtype, sfb_smem)
    ) * atom_thr_size

    size_tensormap_in_i64 = num_tensormaps * bytes_per_tensormap // 8

    @cute.struct
    class LauncherSharedStorage:
        tensormap_buffer: cute.struct.MemRange[cutlass.Int64, size_tensormap_in_i64]
        ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_AB_STAGE]
        ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_AB_STAGE]
        acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_ACC_STAGE]
        acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_ACC_STAGE]
        tmem_dealloc_mbar_ptr: cutlass.Int64
        tmem_holding_buf: cutlass.Int32
        sC: cute.struct.Align[
            cute.struct.MemRange[c_dtype, cute.cosize(c_smem_staged.outer)],
            1024,
        ]
        sA: cute.struct.Align[
            cute.struct.MemRange[ab_dtype, cute.cosize(a_smem_staged.outer)],
            1024,
        ]
        sB: cute.struct.Align[
            cute.struct.MemRange[ab_dtype, cute.cosize(b_smem_staged.outer)],
            1024,
        ]
        sSFA: cute.struct.Align[
            cute.struct.MemRange[sf_dtype, cute.cosize(sfa_smem_staged)],
            1024,
        ]
        sSFB: cute.struct.Align[
            cute.struct.MemRange[sf_dtype, cute.cosize(sfb_smem_staged)],
            1024,
        ]

    cta_tile_shape_mnk = (mma_tiler_mnk[0] // cute.size(tiled_mma.thr_id.shape), mma_tiler_mnk[1], mma_tiler_mnk[2])
    cluster_tile_shape_mnk = tuple(x * y for x, y in zip(cta_tile_shape_mnk, (*cluster_shape_mn, 1)))

    problem_shape_ntile_mnl = (cluster_shape_mn[0], cluster_shape_mn[1], cutlass.Int32(total_num_clusters))
    tile_sched_params = utils.PersistentTileSchedulerParams(problem_shape_ntile_mnl, (*cluster_shape_mn, 1))
    grid = utils.StaticPersistentTileScheduler.get_grid_shape(tile_sched_params, max_active_clusters)

    kernel(
        tiled_mma, tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b,
        tma_atom_sfa, tma_tensor_sfa, tma_atom_sfb, tma_tensor_sfb,
        tma_atom_c, tma_tensor_c,
        a_smem_staged, b_smem_staged, sfa_smem_staged, sfb_smem_staged,
        c_smem_staged, epi_tile,
        tile_sched_params, len(problem_sizes), tensor_problem_sizes,
        tensor_abc, tensor_sfasfb, tensor_tensormap,
        num_tma_load_bytes, NUM_AB_STAGE, NUM_ACC_STAGE, NUM_C_STAGE,
        cluster_tile_shape_mnk, c_layout,
    ).launch(
        grid=grid,
        block=[threads_per_cta, 1, 1],
        cluster=(*cluster_shape_mn, 1),
        smem=LauncherSharedStorage.size_in_bytes(),
        min_blocks_per_mp=1,
    )

_kernel_cache = {}


def _fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
                         abc_cpu, sfasfb_cpu, psizes_cpu, num_groups):
    """Fill pinned CPU buffers with current data pointers and problem sizes."""
    for i in range(num_groups):
        a, b, c = abc_tensors[i]
        sfa, sfb = sfasfb_reordered_tensors[i]
        abc_cpu[i, 0] = a.data_ptr()
        abc_cpu[i, 1] = b.data_ptr()
        abc_cpu[i, 2] = c.data_ptr()
        sfasfb_cpu[i, 0] = sfa.data_ptr()
        sfasfb_cpu[i, 1] = sfb.data_ptr()
        psizes_cpu[i, 0] = problem_sizes[i][0]
        psizes_cpu[i, 1] = problem_sizes[i][1]
        psizes_cpu[i, 2] = problem_sizes[i][2]
        psizes_cpu[i, 3] = problem_sizes[i][3]


def custom_kernel(data: input_t) -> output_t:
    abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
    num_groups = len(problem_sizes)

    cluster_tile_m = mma_tiler_mnk[0] * cluster_shape_mn[0]
    cluster_tile_n = mma_tiler_mnk[1] * cluster_shape_mn[1]
    total_num_clusters = 0
    for m, n, k, l in problem_sizes:
        total_num_clusters += -(-m // cluster_tile_m) * -(-n // cluster_tile_n)
    cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]
    num_tensormap_entries = total_num_clusters * cluster_size

    cache_key = (num_groups, total_num_clusters)

    if cache_key not in _kernel_cache:
        # ---- Compile kernel ----
        null_ptr = make_ptr(cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16)
        null_i64 = make_ptr(cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16)

        hardware_info = cutlass.utils.HardwareInfo()
        max_active_clusters = hardware_info.get_max_active_clusters(cluster_size)

        dummy_sizes = [(128, 128, 256, 1)] * num_groups
        compiled = cute.compile(
            launch_kernel,
            null_ptr, null_i64, null_i64, null_i64,
            total_num_clusters,
            num_tensormap_entries,
            dummy_sizes,
            max_active_clusters,
            options="--opt-level 2",
        )

        # ---- Pre-allocate GPU + pinned CPU buffers ----
        t_problem_sizes = torch.empty((num_groups, 4), dtype=torch.int32, device="cuda")
        t_tensormap = torch.empty(
            (num_tensormap_entries, num_tensormaps, bytes_per_tensormap // 8),
            dtype=torch.int64, device="cuda",
        )
        t_abc_ptrs = torch.empty((num_groups, 3), dtype=torch.int64, device="cuda")
        t_sfasfb_ptrs = torch.empty((num_groups, 2), dtype=torch.int64, device="cuda")

        abc_cpu = torch.empty((num_groups, 3), dtype=torch.int64, pin_memory=True)
        sfasfb_cpu = torch.empty((num_groups, 2), dtype=torch.int64, pin_memory=True)
        psizes_cpu = torch.empty((num_groups, 4), dtype=torch.int32, pin_memory=True)

        p_sizes = make_ptr(cutlass.Int32, t_problem_sizes.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
        p_abc = make_ptr(cutlass.Int64, t_abc_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
        p_sfasfb = make_ptr(cutlass.Int64, t_sfasfb_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
        p_tmap = make_ptr(cutlass.Int64, t_tensormap.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)

        # ---- Warmup call (produces correct results for first invocation) ----
        _fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
                             abc_cpu, sfasfb_cpu, psizes_cpu, num_groups)
        t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
        t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
        t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
        compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
        torch.cuda.synchronize()

        # ---- Try CUDA Graph capture (H2D copies + kernel in one graph) ----
        cuda_graph = None
        try:
            cuda_graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(cuda_graph):
                t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
                t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
                t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
                compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
        except Exception:
            cuda_graph = None

        # Verify graph actually captured the kernel (not just H2D copies)
        if cuda_graph is not None:
            abc_tensors[0][2].fill_(0)
            cuda_graph.replay()
            torch.cuda.synchronize()
            if abc_tensors[0][2].abs().max().item() < 1e-6:
                cuda_graph = None  # Graph didn't capture kernel launch
                # Re-run kernel directly to restore correct C values
                compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
                torch.cuda.synchronize()

        _kernel_cache[cache_key] = (
            compiled, cuda_graph,
            t_abc_ptrs, t_sfasfb_ptrs, t_problem_sizes,
            abc_cpu, sfasfb_cpu, psizes_cpu,
            p_sizes, p_abc, p_sfasfb, p_tmap,
            t_tensormap,
        )

        # First call: warmup (or verification replay) produced correct results
        return [abc_tensors[i][2] for i in range(num_groups)]

    (compiled, cuda_graph,
     t_abc_ptrs, t_sfasfb_ptrs, t_problem_sizes,
     abc_cpu, sfasfb_cpu, psizes_cpu,
     p_sizes, p_abc, p_sfasfb, p_tmap, _) = _kernel_cache[cache_key]

    # Fill pinned CPU buffers (CPU-only, no CUDA calls)
    _fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
                         abc_cpu, sfasfb_cpu, psizes_cpu, num_groups)

    if cuda_graph is not None:
        # CUDA Graph replay: H2D copies + kernel launch in ~3-5μs
        cuda_graph.replay()
    else:
        # Fallback: individual copies + dispatch
        t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
        t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
        t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
        compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)

    return [abc_tensors[i][2] for i in range(num_groups)]
scrolls · 890 lines total

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

Best evidence level for this revision: reported

JSON