Skip to content
KernelIndex
Search⌘K

submission 476474

currybab · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_c.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-476474?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
45.7µs
#51 of 145
2026-02-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d83d68dd288ea4e3e86b814c4fb1c689b2e83a1a025ade04fc2c7bb690c169f0
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-15

Techniques

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

mbarriertmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=threads_per_cta)
shared-memorya_smem_layout_staged: cute.ComposedLayout,
tcgen05HAS_CTA_TWO = hasattr(tcgen05.CtaGroup, "TWO")
tile-k = 256TILE_K = 256
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission_c.py1069 lines
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.cute.runtime import make_ptr

import torch
import traceback
from typing import Dict, Any, Tuple
from task import input_t, output_t

# -----------------------------------------------------------------------------
# Common constants
# -----------------------------------------------------------------------------
bytes_per_tensormap = 128
num_tensormaps = 4

ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16

sf_vec_size = 16
mma_inst_shape_k = 64

# K tile fixed to 256 (task guarantees K%256==0 in bench; tests are multiples of 256 too)
TILE_M_LARGE = 128
TILE_M_SMALL = 64
TILE_K = 256

num_acc_stage = 1
# Stage presets for AB pipeline depth tuning
AB_STAGE_SHORT = 3
AB_STAGE_DEEP = 4

# TMEM allocator constraint: power-of-2, multiple of 32, <=512
num_tmem_alloc_cols = 512

# Keep small-M path compiled-in for future experiments, but disable by default.
# This avoids accidental tile_m=64 selection that caused ms-level regressions.
ENABLE_SMALL_M_EXPERIMENT = False

# Bump when variant/caching behavior changes to avoid mixing stale entries.
CACHE_SCHEMA_VERSION = 3


def ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b


def _raise_picklable(prefix: str):
    tb = traceback.format_exc()
    raise RuntimeError(f"{prefix}\n{tb}") from None


# -----------------------------------------------------------------------------
# Variant selection
# -----------------------------------------------------------------------------
# variant fields:
# - tile_m: 64 or 128
# - tile_n/cta group: narrow(128, ONE) or wide(256, TWO)
# - ab_stage: 3 or 4

HAS_CTA_TWO = hasattr(tcgen05.CtaGroup, "TWO")
HAS_REP_X256 = hasattr(tcgen05.Repetition, "x256")


def _should_use_wide_n(problem_sizes) -> bool:
    """
    Use wide-N only when it is very likely beneficial and safe:
      - all N divisible by 256
      - N is "big enough" (avoid wasting on tiny Ns)
      - and the DSL has the needed enums
    """
    if not (HAS_CTA_TWO and HAS_REP_X256):
        return False
    ns = [int(n) for (_m, n, _k, _l) in problem_sizes]
    if not ns:
        return False
    if max(ns) < 2048:
        return False
    return all((n % 256) == 0 for n in ns)


def _pick_ab_stages(problem_sizes) -> int:
    ks = [int(k) for (_m, _n, k, _l) in problem_sizes]
    if not ks:
        return AB_STAGE_DEEP
    # Keep short pipeline only for very short K; K=2048 tends to prefer deep staging.
    return AB_STAGE_SHORT if max(ks) <= 1536 else AB_STAGE_DEEP


def _should_use_small_m(problem_sizes) -> bool:
    if not ENABLE_SMALL_M_EXPERIMENT:
        return False
    ms = [int(m) for (m, _n, _k, _l) in problem_sizes]
    if not ms:
        return False
    # Small-M grouped shapes benefit from reducing wasted tail compute.
    return max(ms) <= 256


def _make_variant_id(use_small_m: bool, use_wide: bool, ab_stage: int) -> int:
    # bit2: small_m, bit1: wide_n, bit0: deep_stage
    stage_idx = 0 if ab_stage == AB_STAGE_SHORT else 1
    return (4 if use_small_m else 0) + (2 if use_wide else 0) + stage_idx


def _decode_variant_id(variant_id: int) -> Tuple[bool, bool, int]:
    # Returns (use_small_m, use_wide, ab_stage)
    use_small_m = (variant_id & 4) != 0
    use_wide = (variant_id & 2) != 0
    ab_stage = AB_STAGE_SHORT if (variant_id % 2) == 0 else AB_STAGE_DEEP
    return use_small_m, use_wide, ab_stage


# -----------------------------------------------------------------------------
# Kernel/JIT factory per variant (single-kernel launch, tensormap updated per CTA)
# -----------------------------------------------------------------------------
def make_variant(
    tile_m: int,
    tile_n: int,
    cta_group,
    threads_per_cta: int,
    t2r_rep,
    num_ab_stage: int,
):
    mma_tiler_mnk = (tile_m, tile_n, TILE_K)

    @cute.kernel
    def gemm_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,
        tensor_of_abc_ptrs: cute.Tensor,          # [G,3] int64
        tensor_of_sfasfb_ptrs: cute.Tensor,       # [G,2] int64
        tensormaps: cute.Tensor,                  # [T,4,16] int64 (one desc per CTA)
        tensor_of_problem_sizes: cute.Tensor,     # [G,4] int32
        tensor_of_cluster_meta: cute.Tensor,      # [T,3] int32 -> (g, tx, ty)
        a_smem_layout_staged: cute.ComposedLayout,
        b_smem_layout_staged: cute.ComposedLayout,
        sfa_smem_layout_staged: cute.Layout,
        sfb_smem_layout_staged: cute.Layout,
        skip_tensormap_update: cutlass.Int32,
        num_tma_load_bytes: cutlass.Constexpr[int],
    ):
        warp_idx = cute.arch.warp_idx()
        warp_idx = cute.arch.make_warp_uniform(warp_idx)
        tidx, _, _ = cute.arch.thread_idx()

        _, _, bidz = cute.arch.block_idx()

        group_idx = tensor_of_cluster_meta[bidz, 0]
        coord_x = tensor_of_cluster_meta[bidz, 1]
        coord_y = tensor_of_cluster_meta[bidz, 2]

        m = tensor_of_problem_sizes[group_idx, 0]
        n = tensor_of_problem_sizes[group_idx, 1]
        k = tensor_of_problem_sizes[group_idx, 2]
        l = cutlass.Int32(1)

        # ----------------------------
        # Build C tensor for this CTA
        # ----------------------------
        mC_mnl_iter = cute.make_ptr(
            c_dtype, tensor_of_abc_ptrs[group_idx, 2], cute.AddressSpace.gmem
        ).align(32)

        mC_mnl_layout = cute.make_layout(
            (m, n, l),
            stride=(cute.assume(n, 32), 1, cute.assume(m * n, 32)),
        )
        mC_mnl = cute.make_tensor(mC_mnl_iter, mC_mnl_layout)

        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (coord_x, coord_y, 0)
        )

        # ----------------------------
        # Shared storage
        # ----------------------------
        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_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab_stage * 2]
            acc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage * 2]
            tmem_holding_buf: cutlass.Int32

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

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

        # ----------------------------
        # Allocate SMEM tiles
        # ----------------------------
        sA = smem.allocate_tensor(
            element_type=ab_dtype,
            layout=a_smem_layout_staged.outer,
            byte_alignment=128,
            swizzle=a_smem_layout_staged.inner,
        )
        sB = smem.allocate_tensor(
            element_type=ab_dtype,
            layout=b_smem_layout_staged.outer,
            byte_alignment=128,
            swizzle=b_smem_layout_staged.inner,
        )
        sSFA = smem.allocate_tensor(
            element_type=sf_dtype,
            layout=sfa_smem_layout_staged,
            byte_alignment=128,
        )
        sSFB = smem.allocate_tensor(
            element_type=sf_dtype,
            layout=sfb_smem_layout_staged,
            byte_alignment=128,
        )

        # ----------------------------
        # Pipelines
        # ----------------------------
        ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        ab_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, 1)
        ab_producer, ab_consumer = pipeline.PipelineTmaUmma.create(
            barrier_storage=storage.ab_mbar_ptr.data_ptr(),
            num_stages=num_ab_stage,
            producer_group=ab_pipeline_producer_group,
            consumer_group=ab_pipeline_consumer_group,
            tx_count=num_tma_load_bytes,
        ).make_participants()

        acc_producer, acc_consumer = pipeline.PipelineUmmaAsync.create(
            barrier_storage=storage.acc_mbar_ptr.data_ptr(),
            num_stages=num_acc_stage,
            producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
            consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, threads_per_cta),
        ).make_participants()

        # ----------------------------
        # Local tile & MMA partitions
        # ----------------------------
        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)
        )

        thr_mma = tiled_mma.get_slice(tidx)
        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)

        # ----------------------------
        # Update tensormaps (per-CTA) into GMEM buffer, using SMEM scratch
        # ----------------------------
        tensormap_manager = utils.TensorMapManager(utils.TensorMapUpdateMode.SMEM, 128)

        tensormap_a_gmem_ptr = tensormap_manager.get_tensormap_ptr(
            tensormaps[(bidz, 0, None)].iterator
        )
        tensormap_b_gmem_ptr = tensormap_manager.get_tensormap_ptr(
            tensormaps[(bidz, 1, None)].iterator
        )
        tensormap_sfa_gmem_ptr = tensormap_manager.get_tensormap_ptr(
            tensormaps[(bidz, 2, None)].iterator
        )
        tensormap_sfb_gmem_ptr = tensormap_manager.get_tensormap_ptr(
            tensormaps[(bidz, 3, None)].iterator
        )

        # Real tensors
        mA_mkl_iter = cute.make_ptr(
            ab_dtype, tensor_of_abc_ptrs[group_idx, 0], cute.AddressSpace.gmem
        ).align(32)
        mB_nkl_iter = cute.make_ptr(
            ab_dtype, tensor_of_abc_ptrs[group_idx, 1], cute.AddressSpace.gmem
        ).align(32)
        sfa_mkl_iter = cute.make_ptr(
            sf_dtype, tensor_of_sfasfb_ptrs[group_idx, 0], cute.AddressSpace.gmem
        ).align(32)
        sfb_nkl_iter = cute.make_ptr(
            sf_dtype, tensor_of_sfasfb_ptrs[group_idx, 1], cute.AddressSpace.gmem
        ).align(32)

        mA_mkl_layout = cute.make_layout(
            (m, k, l),
            stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
        )
        mB_nkl_layout = cute.make_layout(
            (n, k, l),
            stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
        )

        # Scale factors follow cublas block-scaling-factors layout
        atom_shape = ((32, 4), (sf_vec_size, 4))
        atom_stride = ((16, 4), (0, 1))
        sfa_layout = cute.tile_to_shape(
            cute.make_layout(atom_shape, stride=atom_stride),
            mA_mkl_layout.shape,
            (2, 1, 3),
        )
        sfb_layout = cute.tile_to_shape(
            cute.make_layout(atom_shape, stride=atom_stride),
            mB_nkl_layout.shape,
            (2, 1, 3),
        )

        real_tensor_a = cute.make_tensor(mA_mkl_iter, mA_mkl_layout)
        real_tensor_b = cute.make_tensor(mB_nkl_iter, mB_nkl_layout)
        real_tensor_sfa = cute.make_tensor(sfa_mkl_iter, sfa_layout)
        real_tensor_sfb = cute.make_tensor(sfb_nkl_iter, sfb_layout)

        if warp_idx == 0:
            if skip_tensormap_update == cutlass.Int32(0):
                tensormap_manager.init_tensormap_from_atom(tma_atom_a, tensormap_a_smem_ptr, 0)
                tensormap_manager.init_tensormap_from_atom(tma_atom_b, tensormap_b_smem_ptr, 0)
                tensormap_manager.init_tensormap_from_atom(tma_atom_sfa, tensormap_sfa_smem_ptr, 0)
                tensormap_manager.init_tensormap_from_atom(tma_atom_sfb, tensormap_sfb_smem_ptr, 0)

                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,
                    ),
                    0,
                    (
                        tensormap_a_smem_ptr,
                        tensormap_b_smem_ptr,
                        tensormap_sfa_smem_ptr,
                        tensormap_sfb_smem_ptr,
                    ),
                )

                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)

        cute.arch.barrier()

        # Descriptor ptrs for TMA ops
        tma_desc_a = tensormap_manager.get_tensormap_ptr(
            tensormap_a_gmem_ptr, cute.AddressSpace.generic
        )
        tma_desc_b = tensormap_manager.get_tensormap_ptr(
            tensormap_b_gmem_ptr, cute.AddressSpace.generic
        )
        tma_desc_sfa = tensormap_manager.get_tensormap_ptr(
            tensormap_sfa_gmem_ptr, cute.AddressSpace.generic
        )
        tma_desc_sfb = tensormap_manager.get_tensormap_ptr(
            tensormap_sfb_gmem_ptr, cute.AddressSpace.generic
        )

        # ----------------------------
        # TMA partitions
        # ----------------------------
        tAsA, tAgA = cpasync.tma_partition(
            tma_atom_a,
            0,
            cute.make_layout(1),
            cute.group_modes(sA, 0, 3),
            cute.group_modes(tCgA, 0, 3),
        )
        tBsB, tBgB = cpasync.tma_partition(
            tma_atom_b,
            0,
            cute.make_layout(1),
            cute.group_modes(sB, 0, 3),
            cute.group_modes(tCgB, 0, 3),
        )
        tAsSFA, tAgSFA = cpasync.tma_partition(
            tma_atom_sfa,
            0,
            cute.make_layout(1),
            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,
            0,
            cute.make_layout(1),
            cute.group_modes(sSFB, 0, 3),
            cute.group_modes(tCgSFB, 0, 3),
        )
        tBsSFB = cute.filter_zeros(tBsSFB)
        tBgSFB = cute.filter_zeros(tBgSFB)

        # ----------------------------
        # MMA fragments / TMEM
        # ----------------------------
        tCrA = tiled_mma.make_fragment_A(sA)
        tCrB = tiled_mma.make_fragment_B(sB)

        acc_shape = tiled_mma.partition_shape_C(mma_tiler_mnk[:2])
        tCtAcc_fake = tiled_mma.make_fragment_C(acc_shape)

        tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=threads_per_cta)
        tmem = utils.TmemAllocator(storage.tmem_holding_buf, barrier_for_retrieve=tmem_alloc_barrier)
        tmem.allocate(num_tmem_alloc_cols)
        tmem.wait_for_alloc()

        acc_tmem_ptr = tmem.retrieve_ptr(cutlass.Float32)
        tCtAcc = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

        # SFA/SFB TMEM tensors
        sfa_tmem_ptr = cute.recast_ptr(
            acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc),
            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)
            + 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)

        # S2T copy for SFs (match CTA group)
        copy_atom_s2t = cute.make_copy_atom(
            tcgen05.Cp4x32x128bOp(cta_group),
            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_copy_s2t_sfa = tiled_copy_s2t_sfa.get_slice(0)
        tCsSFA_compact_s2t_ = thr_copy_s2t_sfa.partition_S(tCsSFA_compact)
        tCsSFA_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t_sfa, tCsSFA_compact_s2t_
        )
        tCtSFA_compact_s2t = thr_copy_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_copy_s2t_sfb = tiled_copy_s2t_sfb.get_slice(0)
        tCsSFB_compact_s2t_ = thr_copy_s2t_sfb.partition_S(tCsSFB_compact)
        tCsSFB_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t_sfb, tCsSFB_compact_s2t_
        )
        tCtSFB_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB_compact)

        # K tiles
        k_tile_cnt = k // cutlass.Int32(mma_tiler_mnk[2])

        # Fix coords for this CTA
        tAgA = tAgA[(None, coord_x, None, 0)]
        tBgB = tBgB[(None, coord_y, None, 0)]
        tAgSFA = tAgSFA[(None, coord_x, None, 0)]
        tBgSFB = tBgSFB[(None, coord_y, None, 0)]

        # ----------------------------
        # Main loop (warp0 leader)
        # ----------------------------
        if warp_idx == 0:
            acc_empty = acc_producer.acquire_and_advance()
            tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

            # Prologue prefetch
            for pre_k in range(num_ab_stage):
                if pre_k < k_tile_cnt:
                    ab_empty = ab_producer.acquire_and_advance()
                    st = ab_empty.index

                    cute.copy(
                        tma_atom_a,
                        tAgA[(None, pre_k)],
                        tAsA[(None, st)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_a,
                    )
                    cute.copy(
                        tma_atom_b,
                        tBgB[(None, pre_k)],
                        tBsB[(None, st)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_b,
                    )
                    cute.copy(
                        tma_atom_sfa,
                        tAgSFA[(None, pre_k)],
                        tAsSFA[(None, st)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_sfa,
                    )
                    cute.copy(
                        tma_atom_sfb,
                        tBgSFB[(None, pre_k)],
                        tBsSFB[(None, st)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_sfb,
                    )

            # Steady-state
            for k_tile in cutlass.range(k_tile_cnt):
                ab_full = ab_consumer.wait_and_advance()
                st = ab_full.index

                # S2T SFs
                s2t_stage_coord = (None, None, None, None, st)
                cute.copy(
                    tiled_copy_s2t_sfa,
                    tCsSFA_compact_s2t[s2t_stage_coord],
                    tCtSFA_compact_s2t,
                )
                cute.copy(
                    tiled_copy_s2t_sfb,
                    tCsSFB_compact_s2t[s2t_stage_coord],
                    tCtSFB_compact_s2t,
                )

                # MMA
                num_kblocks = cute.size(tCrA, mode=[2])
                for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                    kblock_coord = (None, None, kblock_idx, st)
                    sf_kblock_coord = (None, None, kblock_idx)

                    tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator)
                    tiled_mma.set(tcgen05.Field.SFB, tCtSFB[sf_kblock_coord].iterator)

                    cute.gemm(
                        tiled_mma,
                        tCtAcc,
                        tCrA[kblock_coord],
                        tCrB[kblock_coord],
                        tCtAcc,
                    )
                    tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

                ab_full.release()

                # Refill
                next_k = k_tile + cutlass.Int32(num_ab_stage)
                if next_k < k_tile_cnt:
                    ab_empty = ab_producer.acquire_and_advance()
                    stp = ab_empty.index

                    cute.copy(
                        tma_atom_a,
                        tAgA[(None, next_k)],
                        tAsA[(None, stp)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_a,
                    )
                    cute.copy(
                        tma_atom_b,
                        tBgB[(None, next_k)],
                        tBsB[(None, stp)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_b,
                    )
                    cute.copy(
                        tma_atom_sfa,
                        tAgSFA[(None, next_k)],
                        tAsSFA[(None, stp)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_sfa,
                    )
                    cute.copy(
                        tma_atom_sfb,
                        tBgSFB[(None, next_k)],
                        tBsSFB[(None, stp)],
                        tma_bar_ptr=ab_empty.barrier,
                        tma_desc_ptr=tma_desc_sfb,
                    )

            acc_empty.commit()

        # ----------------------------
        # Epilogue
        # ----------------------------
        op = tcgen05.Ld32x32bOp(t2r_rep, tcgen05.Pack.NONE)
        copy_atom_t2r = cute.make_copy_atom(op, cutlass.Float32)
        tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tCtAcc[None, 0, 0])
        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)

        tDtAcc = thr_copy_t2r.partition_S(tCtAcc[None, 0, 0])
        tDgC = thr_copy_t2r.partition_D(tCgC[None, 0, 0])

        tDrAcc = cute.make_rmem_tensor(tDgC.shape, cutlass.Float32)
        tDrC = cute.make_rmem_tensor(tDgC.shape, c_dtype)

        tmem.relinquish_alloc_permit()
        acc_full = acc_consumer.wait_and_advance()

        cute.copy(tiled_copy_t2r, tDtAcc, tDrAcc)
        acc_vec = tDrAcc.load()
        tDrC.store(acc_vec.to(c_dtype))

        simt_atom = cute.make_copy_atom(
            cute.nvgpu.CopyUniversalOp(), c_dtype, num_bits_per_copy=16
        )
        thread_layout = cute.make_layout((1, threads_per_cta), stride=(threads_per_cta, 1))
        value_layout = cute.make_layout((1, 1))
        tiled_copy_r2g = cute.make_tiled_copy_tv(simt_atom, thread_layout, value_layout)
        thr_copy_r2g = tiled_copy_r2g.get_slice(tidx)

        residue_m = mC_mnl.shape[0] - cutlass.Int32(coord_x) * cutlass.Int32(mma_tiler_mnk[0])
        residue_n = mC_mnl.shape[1] - cutlass.Int32(coord_y) * cutlass.Int32(mma_tiler_mnk[1])
        # Skip predicate construction on full tiles.
        if residue_m >= cutlass.Int32(mma_tiler_mnk[0]):
            if residue_n >= cutlass.Int32(mma_tiler_mnk[1]):
                cute.copy(simt_atom, cute.flatten(tDrC), cute.flatten(tDgC))
            else:
                cC = cute.make_identity_tensor(gC_mnl.shape)
                tDcC = thr_copy_r2g.partition_D(cC)
                tDpC = cute.make_rmem_tensor(tDrC.shape, cutlass.Boolean)
                for i in range(cute.size(tDrC.shape)):
                    # tDcC ordering is (n,m)
                    tDpC[i] = cute.elem_less(tDcC[i], (residue_n, residue_m))

                cute.copy(
                    simt_atom,
                    cute.flatten(tDrC),
                    cute.flatten(tDgC),
                    pred=cute.flatten(tDpC),
                )
        else:
            cC = cute.make_identity_tensor(gC_mnl.shape)
            tDcC = thr_copy_r2g.partition_D(cC)
            tDpC = cute.make_rmem_tensor(tDrC.shape, cutlass.Boolean)
            for i in range(cute.size(tDrC.shape)):
                # tDcC ordering is (n,m)
                tDpC[i] = cute.elem_less(tDcC[i], (residue_n, residue_m))

            cute.copy(
                simt_atom,
                cute.flatten(tDrC),
                cute.flatten(tDgC),
                pred=cute.flatten(tDpC),
            )

        acc_full.release()
        cute.arch.barrier()
        tmem.free(acc_tmem_ptr)
        pass

    @cute.jit
    def gemm_jit(
        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_cluster_meta: cute.Pointer,
        ptr_of_tensor_of_tensormap: cute.Pointer,
        total_num_clusters: cutlass.Int32,
        num_groups: cutlass.Int32,
        skip_tensormap_update: cutlass.Int32,
    ):
        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))
        )
        tensor_of_problem_sizes = cute.make_tensor(
            ptr_of_tensor_of_problem_sizes, cute.make_layout((num_groups, 4), stride=(4, 1))
        )
        tensor_of_cluster_meta = cute.make_tensor(
            ptr_of_tensor_of_cluster_meta,
            cute.make_layout((total_num_clusters, 3), stride=(3, 1)),
        )
        tensormaps = cute.make_tensor(
            ptr_of_tensor_of_tensormap,
            cute.make_layout((total_num_clusters, 4, 16), stride=(64, 16, 1)),
        )

        # Fake tensors for atom creation: must satisfy M-mode=128 (and N-mode=tile_n)
        min_m = cutlass.Int32(mma_tiler_mnk[0])
        min_n = cutlass.Int32(mma_tiler_mnk[1])
        min_k = cutlass.Int32(mma_tiler_mnk[2])
        one = cutlass.Int32(1)

        initial_a = cute.make_tensor(
            cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16),
            cute.make_layout(
                (min_m, cute.assume(min_k, 32), one),
                stride=(cute.assume(min_k, 32), 1, cute.assume(min_m * min_k, 32)),
            ),
        )
        initial_b = cute.make_tensor(
            cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16),
            cute.make_layout(
                (min_n, cute.assume(min_k, 32), one),
                stride=(cute.assume(min_k, 32), 1, cute.assume(min_n * min_k, 32)),
            ),
        )

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

        mma_op = tcgen05.MmaMXF4NVF4Op(
            sf_dtype,
            (mma_tiler_mnk[0], mma_tiler_mnk[1], mma_inst_shape_k),
            cta_group,
            tcgen05.OperandSource.SMEM,
        )
        tiled_mma = cute.make_tiled_mma(mma_op)

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

        a_smem_layout_staged = sm100_utils.make_smem_layout_a(
            tiled_mma, mma_tiler_mnk, ab_dtype, num_ab_stage
        )
        b_smem_layout_staged = sm100_utils.make_smem_layout_b(
            tiled_mma, mma_tiler_mnk, ab_dtype, num_ab_stage
        )
        sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab_stage
        )
        sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab_stage
        )

        a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, None, 0))
        b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, None, 0))
        sfa_smem_layout = cute.slice_(sfa_smem_layout_staged, (None, None, None, 0))
        sfb_smem_layout = cute.slice_(sfb_smem_layout_staged, (None, None, None, 0))

        # TMA atoms (use matching CTA group)
        tma_kind = cpasync.CopyBulkTensorTileG2SOp(cta_group)

        tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
            tma_kind,
            initial_a,
            a_smem_layout,
            mma_tiler_mnk,
            tiled_mma,
            cluster_layout_vmnk.shape,
        )
        tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
            tma_kind,
            initial_b,
            b_smem_layout,
            mma_tiler_mnk,
            tiled_mma,
            cluster_layout_vmnk.shape,
        )
        tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
            tma_kind,
            initial_sfa,
            sfa_smem_layout,
            mma_tiler_mnk,
            tiled_mma,
            cluster_layout_vmnk.shape,
            internal_type=cutlass.Int16,
        )
        tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
            tma_kind,
            initial_sfb,
            sfb_smem_layout,
            mma_tiler_mnk,
            tiled_mma,
            cluster_layout_vmnk.shape,
            internal_type=cutlass.Int16,
        )

        atom_thr_size = cute.size(tiled_mma.thr_id.shape)

        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)
        num_tma_load_bytes = (a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size) * atom_thr_size

        gemm_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,
            tensor_of_abc_ptrs,
            tensor_of_sfasfb_ptrs,
            tensormaps,
            tensor_of_problem_sizes,
            tensor_of_cluster_meta,
            a_smem_layout_staged,
            b_smem_layout_staged,
            sfa_smem_layout_staged,
            sfb_smem_layout_staged,
            skip_tensormap_update,
            num_tma_load_bytes,
        ).launch(
            grid=(1, 1, total_num_clusters),
            block=[threads_per_cta, 1, 1],
            cluster=(1, 1, 1),
        )
        return

    return gemm_jit, mma_tiler_mnk, threads_per_cta


# -----------------------------------------------------------------------------
# Compile caches per (num_groups, variant_id)
# -----------------------------------------------------------------------------
_compiled_cache: Dict[Tuple[int, int, int], Any] = {}
_compiled_variant_meta: Dict[Tuple[int, int, int], Dict[str, Any]] = {}
_runtime_cache: Dict[Any, Dict[str, Any]] = {}
MAX_SIG_SLOTS = 24


def compile_gemm(num_groups: int, variant_id: int):
    key = (CACHE_SCHEMA_VERSION, num_groups, variant_id)
    if key in _compiled_cache:
        return _compiled_cache[key], _compiled_variant_meta[key]

    use_small_m, use_wide, ab_stage = _decode_variant_id(variant_id)
    tile_m = TILE_M_SMALL if use_small_m else TILE_M_LARGE

    try:
        if use_wide:
            # Wide-N
            gemm_jit, mma_tiler_mnk, threads_per_cta = make_variant(
                tile_m=tile_m,
                tile_n=256,
                cta_group=getattr(tcgen05.CtaGroup, "TWO"),
                threads_per_cta=256,
                t2r_rep=getattr(tcgen05.Repetition, "x256"),
                num_ab_stage=ab_stage,
            )
        else:
            # Narrow
            gemm_jit, mma_tiler_mnk, threads_per_cta = make_variant(
                tile_m=tile_m,
                tile_n=128,
                cta_group=tcgen05.CtaGroup.ONE,
                threads_per_cta=128,
                t2r_rep=tcgen05.Repetition.x128,
                num_ab_stage=ab_stage,
            )

        # Dummy pointers for compilation
        cute_ptr_ps = make_ptr(cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16)
        cute_ptr_abc = make_ptr(cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16)
        cute_ptr_sfs = make_ptr(cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16)
        cute_ptr_meta = make_ptr(cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16)
        cute_ptr_tm = make_ptr(cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16)

        total_clusters = cutlass.Int32(1)
        ng = cutlass.Int32(num_groups)
        skip_update = cutlass.Int32(0)

        fn = cute.compile(
            gemm_jit,
            cute_ptr_ps,
            cute_ptr_abc,
            cute_ptr_sfs,
            cute_ptr_meta,
            cute_ptr_tm,
            total_clusters,
            ng,
            skip_update,
        )

        _compiled_cache[key] = fn
        _compiled_variant_meta[key] = {
            "mma_tiler_mnk": mma_tiler_mnk,
            "threads_per_cta": threads_per_cta,
            "ab_stage": ab_stage,
        }
        return fn, _compiled_variant_meta[key]
    except Exception:
        # Fallback chain:
        # wide -> narrow (same tile_m/stage) -> large_m narrow (same stage) -> large_m narrow deep-stage.
        if use_wide:
            fallback_variant_id = _make_variant_id(use_small_m, False, ab_stage)
            if fallback_variant_id != variant_id:
                return compile_gemm(num_groups, fallback_variant_id)

        if use_small_m:
            fallback_variant_id = _make_variant_id(False, False, ab_stage)
            if fallback_variant_id != variant_id:
                return compile_gemm(num_groups, fallback_variant_id)

        if ab_stage == AB_STAGE_SHORT:
            fallback_variant_id = _make_variant_id(False, False, AB_STAGE_DEEP)
            if fallback_variant_id != variant_id:
                return compile_gemm(num_groups, fallback_variant_id)

        _raise_picklable("compile_gemm failed")


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

        # Choose variant axes: tile_m + wide_n + stage depth
        use_small_m = _should_use_small_m(problem_sizes)
        want_wide = _should_use_wide_n(problem_sizes)
        ab_stage = _pick_ab_stages(problem_sizes)
        variant_id = _make_variant_id(use_small_m, want_wide, ab_stage)

        gemm_fn, vmeta = compile_gemm(num_groups, variant_id)
        mma_tiler_mnk = vmeta["mma_tiler_mnk"]
        tile_n = int(mma_tiler_mnk[1])

        ps_key = tuple((int(m), int(n), int(k), int(l)) for (m, n, k, l) in problem_sizes)
        cache_key = (CACHE_SCHEMA_VERSION, variant_id, num_groups, ps_key)

        if cache_key not in _runtime_cache:
            # Build cluster_meta (ty outer is usually good for B/C contiguity)
            cluster_meta = []
            total_num_clusters = 0
            for g, (m, n, k, l) in enumerate(problem_sizes):
                tiles_m = ceil_div(int(m), int(mma_tiler_mnk[0]))
                tiles_n = ceil_div(int(n), tile_n)
                total_num_clusters += tiles_m * tiles_n
                for ty in range(tiles_n):
                    for tx in range(tiles_m):
                        cluster_meta.append((g, tx, ty))

            tensor_of_problem_sizes = torch.tensor(problem_sizes, dtype=torch.int32, device="cuda")
            tensor_of_cluster_meta = torch.tensor(cluster_meta, dtype=torch.int32, device="cuda")
            host_abc = torch.empty((num_groups, 3), dtype=torch.int64, device="cpu", pin_memory=True)
            host_sfs = torch.empty((num_groups, 2), dtype=torch.int64, device="cpu", pin_memory=True)

            _runtime_cache[cache_key] = {
                "tensor_of_problem_sizes": tensor_of_problem_sizes,
                "tensor_of_cluster_meta": tensor_of_cluster_meta,
                "cute_ptr_ps": make_ptr(
                    cutlass.Int32, tensor_of_problem_sizes.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
                ),
                "cute_ptr_meta": make_ptr(
                    cutlass.Int32, tensor_of_cluster_meta.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
                ),
                "total_num_clusters": total_num_clusters,
                "total_num_clusters_i32": cutlass.Int32(total_num_clusters),
                "num_groups_i32": cutlass.Int32(num_groups),
                "host_abc": host_abc,
                "host_sfs": host_sfs,
                "sig_slots": {},
                "sig_fifo": [],
            }

        st = _runtime_cache[cache_key]

        # Resolve per-signature slot: on cache hit we can skip tensormap update.
        sig = []
        for (a, b, c), (sfa_r, sfb_r), _sz in zip(abc_tensors, sfasfb_reordered_tensors, problem_sizes):
            sig.extend([a.data_ptr(), b.data_ptr(), c.data_ptr(), sfa_r.data_ptr(), sfb_r.data_ptr()])
        sig = tuple(sig)

        sig_slots = st["sig_slots"]
        slot = sig_slots.get(sig)
        if slot is None:
            if len(st["sig_fifo"]) >= MAX_SIG_SLOTS:
                old_sig = st["sig_fifo"].pop(0)
                if old_sig in sig_slots:
                    del sig_slots[old_sig]

            tensor_of_abc_ptrs = torch.empty((num_groups, 3), dtype=torch.int64, device="cuda")
            tensor_of_sfs_ptrs = torch.empty((num_groups, 2), dtype=torch.int64, device="cuda")
            tensor_of_tensormap = torch.empty(
                (st["total_num_clusters"], 4, bytes_per_tensormap // 8),
                dtype=torch.int64,
                device="cuda",
            )
            host_abc = st["host_abc"]
            host_sfs = st["host_sfs"]

            for gi in range(num_groups):
                a, b, c = abc_tensors[gi]
                sfa_r, sfb_r = sfasfb_reordered_tensors[gi]
                host_abc[gi, 0] = a.data_ptr()
                host_abc[gi, 1] = b.data_ptr()
                host_abc[gi, 2] = c.data_ptr()
                host_sfs[gi, 0] = sfa_r.data_ptr()
                host_sfs[gi, 1] = sfb_r.data_ptr()

            tensor_of_abc_ptrs.copy_(host_abc, non_blocking=True)
            tensor_of_sfs_ptrs.copy_(host_sfs, non_blocking=True)

            slot = {
                "tensor_of_abc_ptrs": tensor_of_abc_ptrs,
                "tensor_of_sfs_ptrs": tensor_of_sfs_ptrs,
                "tensor_of_tensormap": tensor_of_tensormap,
                "cute_ptr_abc": make_ptr(
                    cutlass.Int64, tensor_of_abc_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
                ),
                "cute_ptr_sfs": make_ptr(
                    cutlass.Int64, tensor_of_sfs_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
                ),
                "cute_ptr_tm": make_ptr(
                    cutlass.Int64, tensor_of_tensormap.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
                ),
                "tensormap_ready": False,
            }
            sig_slots[sig] = slot
            st["sig_fifo"].append(sig)

        skip_update = cutlass.Int32(1 if slot["tensormap_ready"] else 0)

        # Single launch
        gemm_fn(
            st["cute_ptr_ps"],
            slot["cute_ptr_abc"],
            slot["cute_ptr_sfs"],
            st["cute_ptr_meta"],
            slot["cute_ptr_tm"],
            st["total_num_clusters_i32"],
            st["num_groups_i32"],
            skip_update,
        )
        if not slot["tensormap_ready"]:
            slot["tensormap_ready"] = True

        return [abc_tensors[i][2] for i in range(num_groups)]
    except Exception:
        _raise_picklable("custom_kernel failed")
scrolls · 1069 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 466012.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON