Skip to content
KernelIndex
Search⌘K

submission 202766

Simon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-202766?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
44.9µs
#295 of 420
2025-12-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:25cf04a32ed4874ec70c97db6c952c3fa7d4d951ee9f8ebef589954ef6666725
license declaredunknown
license concludedunknown
authorsSimon
imported2026-08-15

Techniques

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

fused-epilogueepilogue_op: cutlass.Constexpr = lambda x: x
mbarriertmem_alloc_barrier = pipeline.NamedBarrier(
shared-memorya_smem_layout_staged: cute.ComposedLayout,
tcgen05acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),
warp-specializationab1_producer, ab1_consumer = pipeline.PipelineTmaUmma.create(

Kernel source

baseline.py907 lines
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

# Kernel configuration parameters
mma_tiler_mnk = (128, 128, 256)
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  # 4 warps

# Separate stages for each pipeline
num_ab1_stage = 3
num_ab2_stage = 3
num_acc_stage = 1

num_tmem_alloc_cols = 512


@cute.kernel
def kernel(
    tiled_mma: cute.TiledMma,
    # GEMM1: A1, B1, SFA1, SFB1
    tma_atom_a1: cute.CopyAtom,
    mA_mkl1: cute.Tensor,
    tma_atom_b1: cute.CopyAtom,
    mB_nkl1: cute.Tensor,
    tma_atom_sfa1: cute.CopyAtom,
    mSFA_mkl1: cute.Tensor,
    tma_atom_sfb1: cute.CopyAtom,
    mSFB_nkl1: cute.Tensor,
    # GEMM2: A2, B2, SFA2, SFB2
    tma_atom_a2: cute.CopyAtom,
    mA_mkl2: cute.Tensor,
    tma_atom_b2: cute.CopyAtom,
    mB_nkl2: cute.Tensor,
    tma_atom_sfa2: cute.CopyAtom,
    mSFA_mkl2: cute.Tensor,
    tma_atom_sfb2: cute.CopyAtom,
    mSFB_nkl2: cute.Tensor,
    # Output
    mC_mnl: cute.Tensor,
    # Layouts
    a_smem_layout_staged: cute.ComposedLayout,
    b_smem_layout_staged: cute.ComposedLayout,
    sfa_smem_layout_staged: cute.Layout,
    sfb_smem_layout_staged: cute.Layout,
    # Separate byte counts for each pipeline
    num_tma_bytes_gemm1: cutlass.Constexpr[int],
    num_tma_bytes_gemm2: cutlass.Constexpr[int],
    epilogue_op: cutlass.Constexpr = lambda x: x
    * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
    warp_idx = cute.arch.warp_idx()
    warp_idx = cute.arch.make_warp_uniform(warp_idx)
    tidx, _, _ = cute.arch.thread_idx()

    bidx, bidy, bidz = cute.arch.block_idx()
    mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
    cta_coord = (bidx, bidy, bidz)
    mma_tile_coord_mnl = (
        cta_coord[0] // cute.size(tiled_mma.thr_id.shape),
        cta_coord[1],
        cta_coord[2],
    )

    #
    # Shared storage with separate barriers for each pipeline
    #
    @cute.struct
    class SharedStorage:
        ab1_bar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab1_stage * 2]
        ab2_bar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab2_stage * 2]
        acc1_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage * 2]
        acc2_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage * 2]
        tmem_holding_buf: cutlass.Int32

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

    # GEMM1 SMEM: A1, B1, SFA1, SFB1
    sA1 = 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,
    )
    sSFA1 = 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,
    )

    # GEMM2 SMEM: A2, B2, SFA2, SFB2
    sA2 = smem.allocate_tensor(
        element_type=ab_dtype,
        layout=a_smem_layout_staged.outer,
        byte_alignment=128,
        swizzle=a_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,
    )
    sSFA2 = smem.allocate_tensor(
        element_type=sf_dtype,
        layout=sfa_smem_layout_staged,
        byte_alignment=128,
    )
    sSFB2 = smem.allocate_tensor(
        element_type=sf_dtype,
        layout=sfb_smem_layout_staged,
        byte_alignment=128,
    )

    #
    # Four independent pipelines
    #
    # GEMM1 TMA pipeline: Warp 0 produces, Warp 2 consumes
    ab1_producer, ab1_consumer = pipeline.PipelineTmaUmma.create(
        barrier_storage=storage.ab1_bar_ptr.data_ptr(),
        num_stages=num_ab1_stage,
        producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
        consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1),
        tx_count=num_tma_bytes_gemm1,
    ).make_participants()

    # GEMM2 TMA pipeline: Warp 1 produces, Warp 3 consumes
    ab2_producer, ab2_consumer = pipeline.PipelineTmaUmma.create(
        barrier_storage=storage.ab2_bar_ptr.data_ptr(),
        num_stages=num_ab2_stage,
        producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
        consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1),
        tx_count=num_tma_bytes_gemm2,
    ).make_participants()

    # ACC1 pipeline: Warp 2 produces, All threads consume (for epilogue)
    acc1_producer, acc1_consumer = pipeline.PipelineUmmaAsync.create(
        barrier_storage=storage.acc1_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()

    # ACC2 pipeline: Warp 3 produces, All threads consume (for epilogue)
    acc2_producer, acc2_consumer = pipeline.PipelineUmmaAsync.create(
        barrier_storage=storage.acc2_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()

    #
    # Partition global tensors
    #
    # GEMM1
    gA1_mkl = cute.local_tile(
        mA_mkl1, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
    )
    gB1_nkl = cute.local_tile(
        mB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )
    gSFA1_mkl = cute.local_tile(
        mSFA_mkl1, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
    )
    gSFB1_nkl = cute.local_tile(
        mSFB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )

    # GEMM2
    gA2_mkl = cute.local_tile(
        mA_mkl2, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
    )
    gB2_nkl = cute.local_tile(
        mB_nkl2, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )
    gSFA2_mkl = cute.local_tile(
        mSFA_mkl2, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
    )
    gSFB2_nkl = cute.local_tile(
        mSFB_nkl2, 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)
    )
    k_tile_cnt = cute.size(gA1_mkl, mode=[3])

    #
    # MMA partitions
    #
    thr_mma = tiled_mma.get_slice(mma_tile_coord_v)

    # GEMM1
    tCgA1 = thr_mma.partition_A(gA1_mkl)
    tCgB1 = thr_mma.partition_B(gB1_nkl)
    tCgSFA1 = thr_mma.partition_A(gSFA1_mkl)
    tCgSFB1 = thr_mma.partition_B(gSFB1_nkl)

    # GEMM2
    tCgA2 = thr_mma.partition_A(gA2_mkl)
    tCgB2 = thr_mma.partition_B(gB2_nkl)
    tCgSFA2 = thr_mma.partition_A(gSFA2_mkl)
    tCgSFB2 = thr_mma.partition_B(gSFB2_nkl)

    tCgC = thr_mma.partition_C(gC_mnl)

    #
    # TMA partitions for GEMM1
    #
    tAsA1, tAgA1 = cpasync.tma_partition(
        tma_atom_a1,
        0,
        cute.make_layout(1),
        cute.group_modes(sA1, 0, 3),
        cute.group_modes(tCgA1, 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),
    )
    tAsSFA1, tAgSFA1 = cpasync.tma_partition(
        tma_atom_sfa1,
        0,
        cute.make_layout(1),
        cute.group_modes(sSFA1, 0, 3),
        cute.group_modes(tCgSFA1, 0, 3),
    )
    tAsSFA1 = cute.filter_zeros(tAsSFA1)
    tAgSFA1 = cute.filter_zeros(tAgSFA1)
    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),
    )
    tBsSFB1 = cute.filter_zeros(tBsSFB1)
    tBgSFB1 = cute.filter_zeros(tBgSFB1)

    #
    # TMA partitions for GEMM2
    #
    tAsA2, tAgA2 = cpasync.tma_partition(
        tma_atom_a2,
        0,
        cute.make_layout(1),
        cute.group_modes(sA2, 0, 3),
        cute.group_modes(tCgA2, 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),
    )
    tAsSFA2, tAgSFA2 = cpasync.tma_partition(
        tma_atom_sfa2,
        0,
        cute.make_layout(1),
        cute.group_modes(sSFA2, 0, 3),
        cute.group_modes(tCgSFA2, 0, 3),
    )
    tAsSFA2 = cute.filter_zeros(tAsSFA2)
    tAgSFA2 = cute.filter_zeros(tAgSFA2)
    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),
    )
    tBsSFB2 = cute.filter_zeros(tBsSFB2)
    tBgSFB2 = cute.filter_zeros(tBgSFB2)

    #
    # MMA fragments
    #
    tCrA1 = tiled_mma.make_fragment_A(sA1)
    tCrB1 = tiled_mma.make_fragment_B(sB1)
    tCrA2 = tiled_mma.make_fragment_A(sA2)
    tCrB2 = tiled_mma.make_fragment_B(sB2)

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

    #
    # TMEM allocation
    #
    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)
    tCtAcc1 = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
    acc_tmem_ptr1 = cute.recast_ptr(
        acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),
        dtype=cutlass.Float32,
    )
    tCtAcc2 = cute.make_tensor(acc_tmem_ptr1, tCtAcc_fake.layout)

    # SFA/SFB TMEM for GEMM1
    tCtSFA1_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)),
    )
    sfa1_tmem_ptr = cute.recast_ptr(
        acc_tmem_ptr
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc2),
        dtype=sf_dtype,
    )
    tCtSFA1 = cute.make_tensor(sfa1_tmem_ptr, tCtSFA1_layout)

    tCtSFB1_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)),
    )
    sfb1_tmem_ptr = cute.recast_ptr(
        acc_tmem_ptr
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFA1),
        dtype=sf_dtype,
    )
    tCtSFB1 = cute.make_tensor(sfb1_tmem_ptr, tCtSFB1_layout)

    # SFA/SFB TMEM for GEMM2
    sfa2_tmem_ptr = cute.recast_ptr(
        acc_tmem_ptr
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFA1)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFB1),
        dtype=sf_dtype,
    )
    tCtSFA2 = cute.make_tensor(sfa2_tmem_ptr, tCtSFA1_layout)

    sfb2_tmem_ptr = cute.recast_ptr(
        acc_tmem_ptr
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
        + tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFA1)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFB1)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFA2),
        dtype=sf_dtype,
    )
    tCtSFB2 = cute.make_tensor(sfb2_tmem_ptr, tCtSFB1_layout)

    #
    # S2T copy setup for GEMM1
    #
    copy_atom_s2t = cute.make_copy_atom(
        tcgen05.Cp4x32x128bOp(tcgen05.CtaGroup.ONE), sf_dtype
    )

    tCsSFA1_compact = cute.filter_zeros(sSFA1)
    tCtSFA1_compact = cute.filter_zeros(tCtSFA1)
    tiled_copy_s2t_sfa = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFA1_compact)
    thr_copy_s2t_sfa = tiled_copy_s2t_sfa.get_slice(0)
    tCsSFA1_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfa, thr_copy_s2t_sfa.partition_S(tCsSFA1_compact)
    )
    tCtSFA1_compact_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA1_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_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfb, thr_copy_s2t_sfb.partition_S(tCsSFB1_compact)
    )
    tCtSFB1_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB1_compact)

    #
    # S2T copy setup for GEMM2
    #
    tCsSFA2_compact = cute.filter_zeros(sSFA2)
    tCtSFA2_compact = cute.filter_zeros(tCtSFA2)
    tCsSFA2_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfa, thr_copy_s2t_sfa.partition_S(tCsSFA2_compact)
    )
    tCtSFA2_compact_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA2_compact)

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

    #
    # Slice to per-tile coords
    #
    tAgA1 = tAgA1[(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])]
    tAgSFA1 = tAgSFA1[(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])]

    tAgA2 = tAgA2[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
    tBgB2 = tBgB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
    tAgSFA2 = tAgSFA2[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
    tBgSFB2 = tBgSFB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]

    #
    # WARP 0: TMA Producer for GEMM1
    #
    if warp_idx == 0:
        for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
            ab1_empty = ab1_producer.acquire_and_advance()
            cute.copy(
                tma_atom_a1,
                tAgA1[(None, ab1_empty.count)],
                tAsA1[(None, ab1_empty.index)],
                tma_bar_ptr=ab1_empty.barrier,
            )
            cute.copy(
                tma_atom_b1,
                tBgB1[(None, ab1_empty.count)],
                tBsB1[(None, ab1_empty.index)],
                tma_bar_ptr=ab1_empty.barrier,
            )
            cute.copy(
                tma_atom_sfa1,
                tAgSFA1[(None, ab1_empty.count)],
                tAsSFA1[(None, ab1_empty.index)],
                tma_bar_ptr=ab1_empty.barrier,
            )
            cute.copy(
                tma_atom_sfb1,
                tBgSFB1[(None, ab1_empty.count)],
                tBsSFB1[(None, ab1_empty.index)],
                tma_bar_ptr=ab1_empty.barrier,
            )

    #
    # WARP 1: TMA Producer for GEMM2
    #
    if warp_idx == 1:
        for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
            ab2_empty = ab2_producer.acquire_and_advance()
            cute.copy(
                tma_atom_a2,
                tAgA2[(None, ab2_empty.count)],
                tAsA2[(None, ab2_empty.index)],
                tma_bar_ptr=ab2_empty.barrier,
            )
            cute.copy(
                tma_atom_b2,
                tBgB2[(None, ab2_empty.count)],
                tBsB2[(None, ab2_empty.index)],
                tma_bar_ptr=ab2_empty.barrier,
            )
            cute.copy(
                tma_atom_sfa2,
                tAgSFA2[(None, ab2_empty.count)],
                tAsSFA2[(None, ab2_empty.index)],
                tma_bar_ptr=ab2_empty.barrier,
            )
            cute.copy(
                tma_atom_sfb2,
                tBgSFB2[(None, ab2_empty.count)],
                tBsSFB2[(None, ab2_empty.index)],
                tma_bar_ptr=ab2_empty.barrier,
            )

    #
    # WARP 2: MMA Consumer for GEMM1
    #
    if warp_idx == 2:
        acc1_empty = acc1_producer.acquire_and_advance()
        tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

        for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
            ab1_full = ab1_consumer.wait_and_advance()

            # S2T for SFA1, SFB1
            s2t_coord1 = (None, None, None, None, ab1_full.index)
            cute.copy(
                tiled_copy_s2t_sfa, tCsSFA1_compact_s2t[s2t_coord1], tCtSFA1_compact_s2t
            )
            cute.copy(
                tiled_copy_s2t_sfb, tCsSFB1_compact_s2t[s2t_coord1], tCtSFB1_compact_s2t
            )

            # MMA1: ACC1 += A1 @ B1
            num_kblocks = cute.size(tCrA1, mode=[2])
            for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                kblock_coord = (None, None, kblock_idx, ab1_full.index)
                sf_kblock_coord = (None, None, kblock_idx)

                tiled_mma.set(tcgen05.Field.SFA, tCtSFA1[sf_kblock_coord].iterator)
                tiled_mma.set(tcgen05.Field.SFB, tCtSFB1[sf_kblock_coord].iterator)
                cute.gemm(
                    tiled_mma,
                    tCtAcc1,
                    tCrA1[kblock_coord],
                    tCrB1[kblock_coord],
                    tCtAcc1,
                )

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

            ab1_full.release()

        acc1_empty.commit()

    #
    # WARP 3: MMA Consumer for GEMM2
    #
    if warp_idx == 3:
        acc2_empty = acc2_producer.acquire_and_advance()
        tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

        for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
            ab2_full = ab2_consumer.wait_and_advance()

            # S2T for SFA2, SFB2
            s2t_coord2 = (None, None, None, None, ab2_full.index)
            cute.copy(
                tiled_copy_s2t_sfa, tCsSFA2_compact_s2t[s2t_coord2], tCtSFA2_compact_s2t
            )
            cute.copy(
                tiled_copy_s2t_sfb, tCsSFB2_compact_s2t[s2t_coord2], tCtSFB2_compact_s2t
            )

            # MMA2: ACC2 += A2 @ B2
            num_kblocks = cute.size(tCrA2, mode=[2])
            for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                kblock_coord = (None, None, kblock_idx, ab2_full.index)
                sf_kblock_coord = (None, None, kblock_idx)

                tiled_mma.set(tcgen05.Field.SFA, tCtSFA2[sf_kblock_coord].iterator)
                tiled_mma.set(tcgen05.Field.SFB, tCtSFB2[sf_kblock_coord].iterator)
                cute.gemm(
                    tiled_mma,
                    tCtAcc2,
                    tCrA2[kblock_coord],
                    tCrB2[kblock_coord],
                    tCtAcc2,
                )

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

            ab2_full.release()

        acc2_empty.commit()

    #
    # EPILOGUE: All 4 warps participate
    #
    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)]

    acc1_full = acc1_consumer.wait_and_advance()
    cute.copy(tiled_copy_t2r, tTR_tAcc1, tTR_rAcc1)
    acc_vec1 = epilogue_op(tTR_rAcc1.load())
    acc2_full = acc2_consumer.wait_and_advance()
    cute.copy(tiled_copy_t2r, tTR_tAcc2, tTR_rAcc2)
    acc_vec2 = tTR_rAcc2.load()

    acc_vec = acc_vec1 * acc_vec2

    tTR_rC.store(acc_vec.to(c_dtype))
    # Store C to global memory
    cute.copy(simt_atom, tTR_rC, tTR_gC)

    acc1_full.release()
    acc2_full.release()

    # Deallocate TMEM
    cute.arch.barrier()
    tmem.free(acc_tmem_ptr)
    return


@cute.jit
def my_kernel(
    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,
    epilogue_op: cutlass.Constexpr = lambda x: x
    * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
    m, n, k, l = problem_size

    # A tensor (shared by both GEMMs, but we create two TMA descriptors)
    a_tensor = cute.make_tensor(
        a_ptr,
        cute.make_layout(
            (m, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
        ),
    )
    b_tensor1 = cute.make_tensor(
        b1_ptr,
        cute.make_layout(
            (n, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
        ),
    )
    b_tensor2 = cute.make_tensor(
        b2_ptr,
        cute.make_layout(
            (n, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
        ),
    )
    c_tensor = cute.make_tensor(
        c_ptr, cute.make_layout((cute.assume(m, 32), n, l), stride=(n, 1, m * n))
    )

    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(b_tensor1.shape, sf_vec_size)
    sfb_tensor1 = cute.make_tensor(sfb1_ptr, sfb_layout)
    sfb_tensor2 = cute.make_tensor(sfb2_ptr, sfb_layout)

    mma_op = tcgen05.MmaMXF4NVF4Op(
        sf_dtype,
        (mma_tiler_mnk[0], mma_tiler_mnk[1], mma_inst_shape_k),
        tcgen05.CtaGroup.ONE,
        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,)
    )

    # Layouts (same for both GEMMs)
    a_smem_layout_staged = sm100_utils.make_smem_layout_a(
        tiled_mma, mma_tiler_mnk, ab_dtype, num_ab1_stage
    )
    b_smem_layout_staged = sm100_utils.make_smem_layout_b(
        tiled_mma, mma_tiler_mnk, ab_dtype, num_ab1_stage
    )
    sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
        tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab1_stage
    )
    sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
        tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab1_stage
    )

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

    # TMA atoms for GEMM1 (A1, B1, SFA1, SFB1)
    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_atom_a1, tma_tensor_a1 = cute.nvgpu.make_tiled_tma_atom_A(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        a_tensor,
        a_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
    )
    tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        b_tensor1,
        b_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
    )
    tma_atom_sfa1, tma_tensor_sfa1 = cute.nvgpu.make_tiled_tma_atom_A(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        sfa_tensor,
        sfa_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )
    tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        sfb_tensor1,
        sfb_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )

    # TMA atoms for GEMM2 (A2, B2, SFA2, SFB2) - same source tensors, different SMEM targets
    tma_atom_a2, tma_tensor_a2 = cute.nvgpu.make_tiled_tma_atom_A(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        a_tensor,
        a_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
    )
    tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        b_tensor2,
        b_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
    )
    tma_atom_sfa2, tma_tensor_sfa2 = cute.nvgpu.make_tiled_tma_atom_A(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        sfa_tensor,
        sfa_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )
    tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
        cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
        sfb_tensor2,
        sfb_smem_layout,
        mma_tiler_mnk,
        tiled_mma,
        cluster_layout_vmnk.shape,
        internal_type=cutlass.Int16,
    )

    # Byte counts for each pipeline
    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_bytes_gemm1 = (
        a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
    ) * atom_thr_size
    num_tma_bytes_gemm2 = (
        a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
    ) * atom_thr_size

    grid = (
        cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),
        cute.ceil_div(c_tensor.shape[1], mma_tiler_mnk[1]),
        c_tensor.shape[2],
    )

    kernel(
        tiled_mma,
        # GEMM1
        tma_atom_a1,
        tma_tensor_a1,
        tma_atom_b1,
        tma_tensor_b1,
        tma_atom_sfa1,
        tma_tensor_sfa1,
        tma_atom_sfb1,
        tma_tensor_sfb1,
        # GEMM2
        tma_atom_a2,
        tma_tensor_a2,
        tma_atom_b2,
        tma_tensor_b2,
        tma_atom_sfa2,
        tma_tensor_sfa2,
        tma_atom_sfb2,
        tma_tensor_sfb2,
        # Output
        c_tensor,
        # Layouts
        a_smem_layout_staged,
        b_smem_layout_staged,
        sfa_smem_layout_staged,
        sfb_smem_layout_staged,
        # Byte counts
        num_tma_bytes_gemm1,
        num_tma_bytes_gemm2,
        epilogue_op,
    ).launch(
        grid=grid,
        block=[threads_per_cta, 1, 1],
        cluster=(1, 1, 1),
    )
    return


# Global cache for compiled kernel
_compiled_kernel_cache = None


def compile_kernel():
    global _compiled_kernel_cache

    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache

    a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    b1_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    b2_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    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)

    _compiled_kernel_cache = cute.compile(
        my_kernel,
        a_ptr,
        b1_ptr,
        b2_ptr,
        sfa_ptr,
        sfb1_ptr,
        sfb2_ptr,
        c_ptr,
        (0, 0, 0, 0),
    )

    return _compiled_kernel_cache


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data

    compiled_func = compile_kernel()

    _, k, _ = a.shape
    m, n, l = c.shape
    k = k * 2

    a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b1_ptr = make_ptr(ab_dtype, b1.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b2_ptr = make_ptr(ab_dtype, b2.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    sfa_ptr = make_ptr(
        sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
    )
    sfb1_ptr = make_ptr(
        sf_dtype, sfb1_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
    )
    sfb2_ptr = make_ptr(
        sf_dtype, sfb2_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
    )

    compiled_func(
        a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l)
    )

    return c
scrolls · 907 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