Skip to content
KernelIndex
Search⌘K

submission 322877

Phạm Lê Huy Hoàng 2006 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-322877?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 dual GEMMsuite of 4 cases
NVIDIA B200
30.3µs
#247 of 420
2026-01-10

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3a1e6ff3a21c068dd4ef56713108dc712c3fe0b93abfc0a17f1cc13c546c04ed
license declaredunknown
license concludedunknown
authorsPhạm Lê Huy Hoàng 2006
imported2026-08-26

Techniques

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

fused-epilogueepilogue_op: cutlass.Constexpr = (
mbarriertmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=threads_per_cta)
shared-memorya_smem_layout_staged: cute.ComposedLayout,
tcgen05acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),
warp-specializationab_producer, ab_consumer = pipeline.PipelineTmaUmma.create(

Kernel source

submission.py811 lines
# [kernelpipe] AUTOGEN {"DISABLE_2CTA": 0, "DISABLE_64": 0, "K_TILE": 256, "M_TILE": 128, "NUM_AB_STAGE": 3, "NUM_ACC_STAGE": 1, "N_TILE": 64, "SPLIT_N": 2, "TMEM_COLS": 512, "VARIANT_DEFAULT": "128"}
import os
import torch
from task import input_t, output_t

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


# =============================================================================
# Global configuration
# =============================================================================
mma_inst_shape_k = 64

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

sf_vec_size = 16
threads_per_cta = 128

num_acc_stage = 1
num_ab_stage  = 3
# TMEM allocator constraint: <= 512, pow2, multiple of 32
TMEM_COLS_DEFAULT = 512
# =============================================================================
# Helpers
# =============================================================================
def _env_flag(name: str, default: str = "0") -> bool:
    v = os.getenv(name, default).strip().lower()
    return v in ("1", "true", "yes", "y", "on")


def _env_str(name: str, default: str = "") -> str:
    return os.getenv(name, default).strip()


def _env_int(name: str, default: int) -> int:
    try:
        return int(os.getenv(name, str(default)).strip())
    except Exception:
        return default


# =============================================================================
# Kernel factory
#
# split_n:
#   - 1: normal mapping
#   - 2: split-N scheduling using 2 CTAs per (N_TILE * 2) span:
#        split_rank = bidx % 2
#        tile_n = bidy*2 + split_rank
#
# NOTE:
#   This is split-N scheduling WITHOUT launch cluster.
#   We always launch with cluster=(1,1,1) to avoid
#   TMA SMEM-layout vs CTA-Vmap mismatch for blockscaled SFB.
# =============================================================================
def _make_dual_gemm_kernel(mma_tiler_mnk: tuple, split_n: int):
    M_TILE, N_TILE, K_TILE = mma_tiler_mnk
    SPLIT_N = split_n  # constexpr in closure
    SFB_TILER_MNK = (N_TILE, M_TILE, K_TILE)  # scale-factor-B tiler (note M/N swap)

    @cute.kernel
    def _kernel(
        tiled_mma: cute.TiledMma,

        tma_atom_a:  cute.CopyAtom, mA_mkl:   cute.Tensor,
        tma_atom_b1: cute.CopyAtom, mB1_nkl:  cute.Tensor,
        tma_atom_b2: cute.CopyAtom, mB2_nkl:  cute.Tensor,

        tma_atom_sfa:  cute.CopyAtom, mSFA_mkl:  cute.Tensor,
        tma_atom_sfb1: cute.CopyAtom, mSFB1_nkl: cute.Tensor,
        tma_atom_sfb2: cute.CopyAtom, mSFB2_nkl: cute.Tensor,

        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,

        # bytes per CTA per stage
        num_tma_load_bytes: cutlass.Constexpr[int],

        epilogue_op: cutlass.Constexpr = (
            lambda x: x * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True)))
        ),
    ):
        warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx())
        tidx, _, _ = cute.arch.thread_idx()
        bidx, bidy, bidz = cute.arch.block_idx()

        # --- split-N scheduling (no launch cluster) ---
        split_rank  = bidx %SPLIT_N
        bidx_linear = bidx // SPLIT_N

        v_tiles = cute.size(tiled_mma.thr_id.shape)
        mma_tile_coord_v = bidx_linear % v_tiles

        tile_m = bidx_linear // v_tiles
        tile_n = bidy * SPLIT_N + split_rank
        tile_l = bidz
        mma_tile_coord_mnl = (tile_m, tile_n, tile_l)

        @cute.struct
        class SharedStorage:
            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

        # ----------------------------
        # Shared memory allocations
        # ----------------------------
        smem = utils.SmemAllocator()
        storage = smem.allocate(SharedStorage)

        sA = smem.allocate_tensor(
            element_type=ab_dtype,
            layout=a_smem_layout_staged.outer,
            byte_alignment=128,
            swizzle=a_smem_layout_staged.inner,
        )
        sB1 = smem.allocate_tensor(
            element_type=ab_dtype,
            layout=b_smem_layout_staged.outer,
            byte_alignment=128,
            swizzle=b_smem_layout_staged.inner,
        )
        sB2 = 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,
        )
        sSFB1 = smem.allocate_tensor(
            element_type=sf_dtype,
            layout=sfb_smem_layout_staged,
            byte_alignment=128,
        )
        sSFB2 = smem.allocate_tensor(
            element_type=sf_dtype,
            layout=sfb_smem_layout_staged,
            byte_alignment=128,
        )

        # ----------------------------
        # Pipelines
        # ----------------------------
        ab_prod_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        ab_cons_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_prod_group,
            consumer_group=ab_cons_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=ab_prod_group,
            consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, threads_per_cta),
        ).make_participants()

        # ----------------------------
        # Tile global tensors
        # ----------------------------
        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
        )
        gB1_nkl = cute.local_tile(
            mB1_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
        )
        gB2_nkl = cute.local_tile(
            mB2_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)
        )
        gSFB1_nkl = cute.local_tile(
            mSFB1_nkl, cute.slice_(SFB_TILER_MNK, (0, None, None)), (None, None, None)
        )
        gSFB2_nkl = cute.local_tile(
            mSFB2_nkl, cute.slice_(SFB_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)
        )

        k_tile_cnt = cute.size(gA_mkl, mode=[3])

        # ----------------------------
        # MMA slice & partitions
        # ----------------------------
        thr_mma = tiled_mma.get_slice(mma_tile_coord_v)

        tCgA   = thr_mma.partition_A(gA_mkl)
        tCgB1  = thr_mma.partition_B(gB1_nkl)
        tCgB2  = thr_mma.partition_B(gB2_nkl)
        tCgSFA = thr_mma.partition_A(gSFA_mkl)
        tCgSFB1= thr_mma.partition_B(gSFB1_nkl)
        tCgSFB2= thr_mma.partition_B(gSFB2_nkl)
        tCgC   = thr_mma.partition_C(gC_mnl)

        # ----------------------------
        # TMA partitions (always 1CTA vmap)
        # ----------------------------
        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),
        )
        tBsB1, tBgB1 = cpasync.tma_partition(
            tma_atom_b1, 0, cute.make_layout(1),
            cute.group_modes(sB1, 0, 3),
            cute.group_modes(tCgB1, 0, 3),
        )
        tBsB2, tBgB2 = cpasync.tma_partition(
            tma_atom_b2, 0, cute.make_layout(1),
            cute.group_modes(sB2, 0, 3),
            cute.group_modes(tCgB2, 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),
        )
        tBsSFB1, tBgSFB1 = cpasync.tma_partition(
            tma_atom_sfb1, 0, cute.make_layout(1),
            cute.group_modes(sSFB1, 0, 3),
            cute.group_modes(tCgSFB1, 0, 3),
        )
        tBsSFB2, tBgSFB2 = cpasync.tma_partition(
            tma_atom_sfb2, 0, cute.make_layout(1),
            cute.group_modes(sSFB2, 0, 3),
            cute.group_modes(tCgSFB2, 0, 3),
        )

        # SF can have structural zeros; keep safe
        tAsSFA,  tAgSFA  = cute.filter_zeros(tAsSFA),  cute.filter_zeros(tAgSFA)
        tBsSFB1, tBgSFB1 = cute.filter_zeros(tBsSFB1), cute.filter_zeros(tBgSFB1)
        tBsSFB2, tBgSFB2 = cute.filter_zeros(tBsSFB2), cute.filter_zeros(tBgSFB2)

        # ----------------------------
        # MMA fragments
        # ----------------------------
        tCrA  = tiled_mma.make_fragment_A(sA)
        tCrB1 = tiled_mma.make_fragment_B(sB1)
        tCrB2 = tiled_mma.make_fragment_B(sB2)

        acc_shape = tiled_mma.partition_shape_C((M_TILE, N_TILE))
        tCtAcc_fake = tiled_mma.make_fragment_C(acc_shape)

        # ----------------------------
        # TMEM allocation
        # ----------------------------
        tmem_cols = _env_int("CUTLASS_TMEM_COLS", TMEM_COLS_DEFAULT)

        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(tmem_cols)
        tmem.wait_for_alloc()

        acc_tmem_ptr = tmem.retrieve_ptr(cutlass.Float32)

        tCtAcc1 = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
        acc2_ptr = cute.recast_ptr(
            acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),
            dtype=cutlass.Float32,
        )
        tCtAcc2 = cute.make_tensor(acc2_ptr, tCtAcc_fake.layout)

        # ----------------------------
        # SFA/SFB in TMEM (offset arithmetic on Float32 base pointer)
        # ----------------------------
        base_f32 = acc_tmem_ptr
        acc1_off = tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
        acc2_off = tcgen05.find_tmem_tensor_col_offset(tCtAcc2)

        # ---- SFA ----
        sfa_base_f32 = base_f32 + acc1_off + acc2_off
        sfa_tmem_ptr = cute.recast_ptr(sfa_base_f32, 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)
        sfa_off = tcgen05.find_tmem_tensor_col_offset(tCtSFA)

        # ---- SFB1 ----
        sfb1_base_f32 = base_f32 + acc1_off + acc2_off + sfa_off
        sfb1_tmem_ptr = cute.recast_ptr(sfb1_base_f32, dtype=sf_dtype)
        

        tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
            tiled_mma, SFB_TILER_MNK, sf_vec_size,
            cute.slice_(sfb_smem_layout_staged, (None, None, None, 0))
        )
        tCtSFB1 = cute.make_tensor(sfb1_tmem_ptr, tCtSFB_layout)
        sfb1_off = tcgen05.find_tmem_tensor_col_offset(tCtSFB1)

        # ---- SFB2 ----
        sfb2_base_f32 = base_f32 + acc1_off + acc2_off + sfa_off + sfb1_off
        sfb2_tmem_ptr = cute.recast_ptr(sfb2_base_f32, dtype=sf_dtype)
        tCtSFB2 = cute.make_tensor(sfb2_tmem_ptr, tCtSFB_layout)

        # ----------------------------
        # S2T copy setup
        # ----------------------------
        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_copy_s2t_sfa   = tiled_copy_s2t_sfa.get_slice(0)
        tCsSFA_desc = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t_sfa, thr_copy_s2t_sfa.partition_S(tCsSFA_compact)
        )
        tCtSFA_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA_compact)

        tCsSFB1_compact = cute.filter_zeros(sSFB1)
        tCtSFB1_compact = cute.filter_zeros(tCtSFB1)
        tiled_copy_s2t_sfb = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFB1_compact)
        thr_copy_s2t_sfb   = tiled_copy_s2t_sfb.get_slice(0)
        tCsSFB1_desc = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t_sfb, thr_copy_s2t_sfb.partition_S(tCsSFB1_compact)
        )
        tCtSFB1_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB1_compact)

        tCsSFB2_compact = cute.filter_zeros(sSFB2)
        tCtSFB2_compact = cute.filter_zeros(tCtSFB2)
        tCsSFB2_desc = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t_sfb, thr_copy_s2t_sfb.partition_S(tCsSFB2_compact)
        )
        tCtSFB2_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB2_compact)

        # ----------------------------
        # Slice gmem tiles to this CTA tile coord
        # ----------------------------
        tAgA   = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
        tBgB1  = tBgB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
        tBgB2  = tBgB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]

        tAgSFA = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
        tBgSFB1= tBgSFB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
        tBgSFB2= tBgSFB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]

        # ----------------------------
        # Mainloop (warp-specialized producer/compute)
        # ----------------------------
        if warp_idx == 0:
            acc_empty = acc_producer.acquire_and_advance()
            tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

            # Prologue
            for k_pro in range(min(k_tile_cnt, num_ab_stage - 1)):
                ab_empty = ab_producer.acquire_and_advance()
                cute.copy(tma_atom_a,    tAgA[(None, k_pro)],     tAsA[(None, ab_empty.index)],     tma_bar_ptr=ab_empty.barrier)
                cute.copy(tma_atom_b1,   tBgB1[(None, k_pro)],    tBsB1[(None, ab_empty.index)],    tma_bar_ptr=ab_empty.barrier)
                cute.copy(tma_atom_b2,   tBgB2[(None, k_pro)],    tBsB2[(None, ab_empty.index)],    tma_bar_ptr=ab_empty.barrier)
                cute.copy(tma_atom_sfa,  tAgSFA[(None, k_pro)],   tAsSFA[(None, ab_empty.index)],   tma_bar_ptr=ab_empty.barrier)
                cute.copy(tma_atom_sfb1, tBgSFB1[(None, k_pro)],  tBsSFB1[(None, ab_empty.index)],  tma_bar_ptr=ab_empty.barrier)
                cute.copy(tma_atom_sfb2, tBgSFB2[(None, k_pro)],  tBsSFB2[(None, ab_empty.index)],  tma_bar_ptr=ab_empty.barrier)

            for k_tile in range(k_tile_cnt):
                if k_tile + num_ab_stage - 1 < k_tile_cnt:
                    k_next = k_tile + num_ab_stage - 1
                    ab_empty = ab_producer.acquire_and_advance()
                    cute.copy(tma_atom_a,    tAgA[(None, k_next)],     tAsA[(None, ab_empty.index)],     tma_bar_ptr=ab_empty.barrier)
                    cute.copy(tma_atom_b1,   tBgB1[(None, k_next)],    tBsB1[(None, ab_empty.index)],    tma_bar_ptr=ab_empty.barrier)
                    cute.copy(tma_atom_b2,   tBgB2[(None, k_next)],    tBsB2[(None, ab_empty.index)],    tma_bar_ptr=ab_empty.barrier)
                    cute.copy(tma_atom_sfa,  tAgSFA[(None, k_next)],   tAsSFA[(None, ab_empty.index)],   tma_bar_ptr=ab_empty.barrier)
                    cute.copy(tma_atom_sfb1, tBgSFB1[(None, k_next)],  tBsSFB1[(None, ab_empty.index)],  tma_bar_ptr=ab_empty.barrier)
                    cute.copy(tma_atom_sfb2, tBgSFB2[(None, k_next)],  tBsSFB2[(None, ab_empty.index)],  tma_bar_ptr=ab_empty.barrier)

                ab_full = ab_consumer.wait_and_advance()

                # S2T for scale factors
                s2t_stage_coord = (None, None, None, None, ab_full.index)
                cute.copy(tiled_copy_s2t_sfa, tCsSFA_desc[s2t_stage_coord],  tCtSFA_s2t)
                cute.copy(tiled_copy_s2t_sfb, tCsSFB1_desc[s2t_stage_coord], tCtSFB1_s2t)
                cute.copy(tiled_copy_s2t_sfb, tCsSFB2_desc[s2t_stage_coord], tCtSFB2_s2t)

                # GEMM
                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_full.index)
                    sf_kblock_coord = (None, None, kblock_idx)

                    # ACC1 += A@B1
                    tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator)
                    tiled_mma.set(tcgen05.Field.SFB, tCtSFB1[sf_kblock_coord].iterator)
                    cute.gemm(tiled_mma, tCtAcc1, tCrA[kblock_coord], tCrB1[kblock_coord], tCtAcc1)

                    # ACC2 += A@B2
                    tiled_mma.set(tcgen05.Field.SFB, tCtSFB2[sf_kblock_coord].iterator)
                    cute.gemm(tiled_mma, tCtAcc2, tCrA[kblock_coord], tCrB2[kblock_coord], tCtAcc2)

                    tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

                ab_full.release()

            acc_empty.commit()

        # ----------------------------
        # Epilogue
        # ----------------------------
        op = tcgen05.Ld32x32bOp(tcgen05.Repetition.x128, tcgen05.Pack.NONE)
        copy_atom_t2r = cute.make_copy_atom(op, cutlass.Float32)
        tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tCtAcc1)

        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
        tTR_tAcc1 = thr_copy_t2r.partition_S(tCtAcc1)
        tTR_tAcc2 = thr_copy_t2r.partition_S(tCtAcc2)
        tTR_gC    = thr_copy_t2r.partition_D(tCgC)

        tTR_rAcc1 = cute.make_rmem_tensor(tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32)
        tTR_rAcc2 = cute.make_rmem_tensor(tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32)
        tTR_rC    = cute.make_rmem_tensor(tTR_gC[None, None, None, None, 0, 0, 0].shape, c_dtype)

        simt_atom = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), c_dtype)
        tTR_gC = tTR_gC[(None, None, None, None, *mma_tile_coord_mnl)]

        acc_full = acc_consumer.wait_and_advance()

        cute.copy(tiled_copy_t2r, tTR_tAcc1, tTR_rAcc1)
        cute.copy(tiled_copy_t2r, tTR_tAcc2, tTR_rAcc2)
        cute.arch.fence_view_async_tmem_load()

        acc1 = epilogue_op(tTR_rAcc1.load())
        acc2 = tTR_rAcc2.load()
        out  = (acc1 * acc2).to(c_dtype)

        tTR_rC.store(out)
        cute.copy(simt_atom, tTR_rC, tTR_gC)

        acc_full.release()

        cute.arch.barrier()
        tmem.free(acc_tmem_ptr)

    return _kernel


# =============================================================================
# Build kernels
# =============================================================================
KERNEL_128x128_1CTA = _make_dual_gemm_kernel((128, 128, 256), split_n=1)
KERNEL_128x64_1CTA  = _make_dual_gemm_kernel((128, 64, 256), split_n=1)
KERNEL_128x64_2CTA  = _make_dual_gemm_kernel((128, 64, 256), split_n=2)  # split-N scheduling


# =============================================================================
# JIT launchers
# =============================================================================
def _make_jit(kernel_fn, mma_tiler_mnk: tuple, split_n: int):
    M_TILE, N_TILE, K_TILE = mma_tiler_mnk
    SPLIT_N = split_n  # constexpr

    @cute.jit
    def _jit(
        a_ptr:    cute.Pointer,
        b1_ptr:   cute.Pointer,
        b2_ptr:   cute.Pointer,
        sfa_ptr:  cute.Pointer,
        sfb1_ptr: cute.Pointer,
        sfb2_ptr: cute.Pointer,
        c_ptr:    cute.Pointer,
        problem_size: tuple,
    ):
        m, n, k, l = problem_size

        m_s = cute.assume(m, 32)
        n_s = cute.assume(n, 32)
        k_s = cute.assume(k, 32)
        l_s = cute.assume(l, 1)

        # L-major physical layout
        a_tensor = cute.make_tensor(
            a_ptr,
            cute.make_layout((m_s, k_s, l_s), stride=(k_s, 1, cute.assume(m * k, 32))),
        )
        b1_tensor = cute.make_tensor(
            b1_ptr,
            cute.make_layout((n_s, k_s, l_s), stride=(k_s, 1, cute.assume(n * k, 32))),
        )
        b2_tensor = cute.make_tensor(b2_ptr, b1_tensor.layout)

        c_tensor = cute.make_tensor(
            c_ptr,
            cute.make_layout((m_s, n_s, l_s), stride=(n_s, 1, cute.assume(m * n, 32))),
        )

        sfa_layout  = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
        sfa_tensor  = cute.make_tensor(sfa_ptr, sfa_layout)

        sfb_layout  = blockscaled_utils.tile_atom_to_shape_SF(b1_tensor.shape, sf_vec_size)
        sfb1_tensor = cute.make_tensor(sfb1_ptr, sfb_layout)
        sfb2_tensor = cute.make_tensor(sfb2_ptr, sfb_layout)

        # MMA
        mma_op = tcgen05.MmaMXF4NVF4Op(
            sf_dtype,
            (M_TILE, N_TILE, mma_inst_shape_k),
            tcgen05.CtaGroup.ONE,
            tcgen05.OperandSource.SMEM,
        )
        tiled_mma = cute.make_tiled_mma(mma_op)

        # ------------------------------------------------------------
        # 1) Build SMEM layouts FIRST
        # ------------------------------------------------------------
        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_tiler_mnk = (N_TILE, M_TILE, K_TILE)

        sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma, sfb_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))

        # ------------------------------------------------------------
        # 2) Build TMA atoms ONCE, force 1CTA vmap
        # ------------------------------------------------------------
        TMA_CLUSTER_SHAPE = (1, 1, 1)

        tma_atom_a,  tma_tensor_a  = cute.nvgpu.make_tiled_tma_atom_A(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            a_tensor, a_smem_layout, mma_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
        )
        tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            b1_tensor, b_smem_layout, mma_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
        )
        tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            b2_tensor, b_smem_layout, mma_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
        )

        tma_atom_sfa,  tma_tensor_sfa  = cute.nvgpu.make_tiled_tma_atom_A(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            sfa_tensor, sfa_smem_layout, mma_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
            internal_type=cutlass.Int16,
        )
        tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            sfb1_tensor, sfb_smem_layout,
            sfb_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
            internal_type=cutlass.Int16,
        )
        tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
            cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
            sfb2_tensor, sfb_smem_layout,
            sfb_tiler_mnk, tiled_mma,
            TMA_CLUSTER_SHAPE,
            internal_type=cutlass.Int16,
        )

        # bytes per stage per CTA
        a_bytes   = cute.size_in_bytes(ab_dtype, a_smem_layout)
        b_bytes   = cute.size_in_bytes(ab_dtype, b_smem_layout)
        sfa_bytes = cute.size_in_bytes(sf_dtype, sfa_smem_layout)
        sfb_bytes = cute.size_in_bytes(sf_dtype, sfb_smem_layout)
        num_tma_load_bytes = (a_bytes + 2 * b_bytes + sfa_bytes + 2 * sfb_bytes)

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

        # grid:
        # x: m_tiles * v_tiles * SPLIT_N  (because kernel uses bidx%SPLIT_N)
        # y: clusters over N, each cluster covers N_TILE * SPLIT_N columns
        grid_x = cute.ceil_div(c_tensor.shape[0], M_TILE) * v_tiles * SPLIT_N
        grid_y = cute.ceil_div(c_tensor.shape[1], N_TILE * SPLIT_N)
        grid   = (grid_x, grid_y, c_tensor.shape[2])

        kernel_fn(
            tiled_mma,
            tma_atom_a,  tma_tensor_a,
            tma_atom_b1, tma_tensor_b1,
            tma_atom_b2, tma_tensor_b2,
            tma_atom_sfa,  tma_tensor_sfa,
            tma_atom_sfb1, tma_tensor_sfb1,
            tma_atom_sfb2, tma_tensor_sfb2,
            c_tensor,
            a_smem_layout_staged, b_smem_layout_staged,
            sfa_smem_layout_staged, sfb_smem_layout_staged,
            num_tma_load_bytes,
        ).launch(
            grid=grid,
            block=[threads_per_cta, 1, 1],
            cluster=(1, 1, 1),  # ✅ FIX: ALWAYS (1,1,1). Split-N is done via bidx mapping.
        )

    return _jit


JIT_128x128_1CTA = _make_jit(KERNEL_128x128_1CTA, (128, 128, 256), split_n=1)
JIT_128x64_1CTA  = _make_jit(KERNEL_128x64_1CTA, (128,  64, 256), split_n=1)
JIT_128x64_2CTA  = _make_jit(KERNEL_128x64_2CTA, (128,  64, 256), split_n=2)


# =============================================================================
# Compile caches (lazy + safe)
# =============================================================================
_compiled_128_1cta = None
_compiled_64_1cta  = None
_compiled_64_2cta  = None

_failed_64_1cta = False
_failed_64_2cta = False


def _compile_jit_once(jit_fn, shape_hint: tuple):
    # Dummy pointers
    a_ptr  = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=128)
    b1_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=128)
    b2_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=128)
    c_ptr  = make_ptr(c_dtype,  0, cute.AddressSpace.gmem, assumed_align=128)

    sfa_ptr  = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
    sfb1_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
    sfb2_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)

    return cute.compile(
        jit_fn,
        a_ptr, b1_ptr, b2_ptr,
        sfa_ptr, sfb1_ptr, sfb2_ptr,
        c_ptr,
        shape_hint,
    )


def _compile_once_128_1cta():
    global _compiled_128_1cta
    if _compiled_128_1cta is None:
        _compiled_128_1cta = _compile_jit_once(JIT_128x128_1CTA, (256, 256, 256, 1))
    return _compiled_128_1cta


def _compile_once_64_1cta():
    global _compiled_64_1cta, _failed_64_1cta
    if _compiled_64_1cta is not None:
        return _compiled_64_1cta
    if _failed_64_1cta:
        return None
    try:
        _compiled_64_1cta = _compile_jit_once(JIT_128x64_1CTA, (256, 256, 256, 1))
        return _compiled_64_1cta
    except Exception:
        _failed_64_1cta = True
        return None


def _compile_once_64_2cta():
    global _compiled_64_2cta, _failed_64_2cta
    if _compiled_64_2cta is not None:
        return _compiled_64_2cta
    if _failed_64_2cta:
        return None
    try:
        _compiled_64_2cta = _compile_jit_once(JIT_128x64_2CTA, (256, 256, 256, 1))
        return _compiled_64_2cta
    except Exception:
        _failed_64_2cta = True
        return None


# =============================================================================
# Variant selection
# =============================================================================
def _pick_variant(m: int, n: int, k: int, l: int) -> str:
    """
    Returns: "128_1cta" | "64_2cta" | "64_1cta"

    Env overrides:
      CUTLASS_VARIANT = auto | 128_1cta | 64_1cta | 64_2cta
      CUTLASS_DISABLE_2CTA=1  -> disable 64_2cta
      CUTLASS_DISABLE_64=1    -> disable 64_1cta (also affects 64_2cta fallback)
    """
    variant = _env_str("CUTLASS_VARIANT", "128").lower()
    disable_2cta = _env_flag("CUTLASS_DISABLE_2CTA", "0")
    disable_64   = _env_flag("CUTLASS_DISABLE_64", "0")

    # Manual overrides
    if variant in ("128", "128_1cta", "baseline", "128x128"):
        return "128_1cta"
    if variant in ("64", "64_1cta", "128x64"):
        return "128_1cta" if disable_64 else "64_1cta"
    if variant in ("2cta", "64_2cta", "split2"):
        if disable_2cta:
            return "64_1cta" if not disable_64 else "128_1cta"
        return "64_2cta"

    # AUTO (ưu tiên 2CTA split-N scheduling khi tile đẹp)
    full_tiles_ok = (l == 1) and (m % 128 == 0) and (k % 256 == 0)

    if full_tiles_ok and (not disable_2cta) and (n % 128 == 0):
        return "64_2cta"

    if full_tiles_ok and (not disable_64) and (n % 64 == 0):
        return "64_1cta"

    return "128_1cta"


# =============================================================================
# Entry point
# =============================================================================
def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_perm, sfb1_perm, sfb2_perm, c = data

    # Torch uses e2m1_x2 (packed), so logical K is doubled
    _, k_half, _ = a.shape
    m, n, l = c.shape
    k = k_half * 2

    a_ptr  = make_ptr(ab_dtype, a.data_ptr(),  cute.AddressSpace.gmem, assumed_align=128)
    b1_ptr = make_ptr(ab_dtype, b1.data_ptr(), cute.AddressSpace.gmem, assumed_align=128)
    b2_ptr = make_ptr(ab_dtype, b2.data_ptr(), cute.AddressSpace.gmem, assumed_align=128)
    c_ptr  = make_ptr(c_dtype,  c.data_ptr(),  cute.AddressSpace.gmem, assumed_align=128)

    sfa_ptr  = make_ptr(sf_dtype, sfa_perm.data_ptr(),  cute.AddressSpace.gmem, assumed_align=32)
    sfb1_ptr = make_ptr(sf_dtype, sfb1_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
    sfb2_ptr = make_ptr(sf_dtype, sfb2_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)

    choice = _pick_variant(m, n, k, l)

    debug = _env_flag("CUTLASS_DEBUG_VARIANT", "0")
    if debug:
        print(f"[cutlass] pick={choice}  (m,n,k,l)=({m},{n},{k},{l})  TMEM_COLS={_env_int('CUTLASS_TMEM_COLS', TMEM_COLS_DEFAULT)}")

    # Try selected variant first; fallback chain
    if choice == "64_2cta":
        compiled = _compile_once_64_2cta()
        if compiled is not None:
            try:
                compiled(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l))
                if debug:
                    print("[cutlass] ran=64_2cta")
                return c
            except Exception:
                if debug:
                    print("[cutlass] ran=64_2cta FAILED -> fallback")

        compiled = _compile_once_64_1cta()
        if compiled is not None:
            try:
                compiled(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l))
                if debug:
                    print("[cutlass] ran=64_1cta (fallback)")
                return c
            except Exception:
                if debug:
                    print("[cutlass] ran=64_1cta FAILED -> fallback")

    if choice == "64_1cta":
        compiled = _compile_once_64_1cta()
        if compiled is not None:
            try:
                compiled(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l))
                if debug:
                    print("[cutlass] ran=64_1cta")
                return c
            except Exception:
                if debug:
                    print("[cutlass] ran=64_1cta FAILED -> fallback")

    compiled = _compile_once_128_1cta()
    compiled(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l))
    if debug:
        print("[cutlass] ran=128_1cta")
    return c
scrolls · 811 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