Skip to content
KernelIndex
Search⌘K

submission 384880

mklz · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

reference.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-384880?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
70.0µs
#159 of 161
2026-01-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3b899c7143824938809f75708616c0d3a6d5c07ccaf1d3c3f13425dadc6356d2
license declaredunknown
license concludedunknown
authorsmklz
imported2026-08-26

Techniques

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

fp4GPU device kernel performing dual block-scaled NVFP4 GEMM with SiLU-gated fusion.
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-specialization- Warp 0 handles all TMA loads and MMA operations (producer)

Kernel source

reference.py1103 lines
from torch._higher_order_ops.torchbind import call_torchbind_fake
import cuda.bindings.driver as cuda

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.torch as cutlass_torch
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
# Tile sizes for M, N, K dimensions
mma_tiler_mnk= (128, 128, 256)  
# Shape of the K dimension for the MMA instruction
mma_inst_shape_k = 64
# FP4 data type for A and B
ab_dtype = cutlass.Float4E2M1FN  
# FP8 data type for scale factors
sf_dtype = cutlass.Float8E4M3FN  
# FP16 output type
c_dtype = cutlass.Float16  
# Scale factor block size (16 elements share one scale)
sf_vec_size = 16  
# Number of threads per CUDA thread block
threads_per_cta = 128  
# Stage numbers of shared memory and tmem
num_acc_stage = 1
num_ab_stage = 1
# Total number of columns in tmem
num_tmem_alloc_cols = 512


# Helper function for ceiling division
def ceil_div(a, b):
    """
    Compute the ceiling of a / b using integer arithmetic.
    
    This is used throughout the kernel to calculate grid dimensions and tile counts
    when the problem size doesn't evenly divide by the tile size.
    
    Args:
        a: Numerator (e.g., problem dimension M, N, or K)
        b: Denominator (e.g., tile size)
    
    Returns:
        ceil(a / b) as an integer
    
    Example:
        ceil_div(1000, 128) = 8  # Need 8 tiles of size 128 to cover 1000 elements
    """
    return (a + b - 1) // b


#  GPU device kernel
@cute.kernel
def kernel(
    tiled_mma: cute.TiledMma,
    tma_atom_a: cute.CopyAtom,
    mA_mkl: cute.Tensor,
    tma_atom_b1: cute.CopyAtom,
    mB_nkl1: cute.Tensor,
    tma_atom_b2: cute.CopyAtom,
    mB_nkl2: cute.Tensor,
    tma_atom_sfa: cute.CopyAtom,
    mSFA_mkl: cute.Tensor,
    tma_atom_sfb1: cute.CopyAtom,
    mSFB_nkl1: cute.Tensor,
    tma_atom_sfb2: cute.CopyAtom,
    mSFB_nkl2: 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,
    num_tma_load_bytes: cutlass.Constexpr[int],
    epilogue_op: cutlass.Constexpr = lambda x: x
    * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
    """
    GPU device kernel performing dual block-scaled NVFP4 GEMM with SiLU-gated fusion.
    
    Computes: C = silu(A @ B1) * (A @ B2)
    
    This kernel executes the following operations:
    1. Loads FP4 matrices A, B1, B2 and their FP8 block scale factors from global memory
       to shared memory using TMA (Tensor Memory Accelerator)
    2. Copies scale factors from shared memory to tensor memory (TMEM)
    3. Performs two block-scaled matrix multiplications using Blackwell's tcgen05 MMA:
       - ACC1 = A @ B1 (with block scaling from SFA and SFB1)
       - ACC2 = A @ B2 (with block scaling from SFA and SFB2)
    4. Applies SiLU activation to ACC1 and multiplies element-wise with ACC2
    5. Stores the FP16 result to global memory
    
    Key optimization: Matrix A and its scale factors (SFA) are loaded once but reused
    for both GEMMs, reducing memory bandwidth by ~33%.
    
    Execution model:
    - Warp 0 handles all TMA loads and MMA operations (producer)
    - All warps participate in the epilogue (copying results to global memory)
    - Pipeline barriers synchronize TMA loads with MMA consumption
    
    Args:
        tiled_mma: TiledMMA object defining the MMA instruction configuration
                   (tile shape, data types, tensor core operation)
        tma_atom_a: TMA copy atom for loading matrix A from global to shared memory
        mA_mkl: Global memory tensor for A with shape (M, K, L) in FP4 format
        tma_atom_b1: TMA copy atom for loading matrix B1
        mB_nkl1: Global memory tensor for B1 with shape (N, K, L) in FP4 format
        tma_atom_b2: TMA copy atom for loading matrix B2
        mB_nkl2: Global memory tensor for B2 with shape (N, K, L) in FP4 format
        tma_atom_sfa: TMA copy atom for loading scale factors for A
        mSFA_mkl: Global memory tensor for SFA (block scale factors for A)
        tma_atom_sfb1: TMA copy atom for loading scale factors for B1
        mSFB_nkl1: Global memory tensor for SFB1 (block scale factors for B1)
        tma_atom_sfb2: TMA copy atom for loading scale factors for B2
        mSFB_nkl2: Global memory tensor for SFB2 (block scale factors for B2)
        mC_mnl: Global memory tensor for output C with shape (M, N, L) in FP16
        a_smem_layout_staged: Shared memory layout for A with staging dimension
        b_smem_layout_staged: Shared memory layout for B1/B2 with staging
        sfa_smem_layout_staged: Shared memory layout for SFA with staging
        sfb_smem_layout_staged: Shared memory layout for SFB1/SFB2 with staging
        num_tma_load_bytes: Total bytes per TMA transaction (for barrier setup)
        epilogue_op: Element-wise operation applied to ACC1 (default: SiLU activation)
    """
    warp_idx = cute.arch.warp_idx()
    warp_idx = cute.arch.make_warp_uniform(warp_idx)
    tidx = cute.arch.thread_idx()

    #
    # Setup cta/thread coordinates
    #
    # Coords inside cluster
    bidx, bidy, bidz = cute.arch.block_idx()
    mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)

    # Coords outside cluster
    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],
    )
    # Coord inside cta
    tidx, _, _ = cute.arch.thread_idx()

    #
    # Define shared storage for kernel
    #
    @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

    smem = utils.SmemAllocator()
    storage = smem.allocate(SharedStorage)
    # (MMA, MMA_M, MMA_K, STAGE)
    sA = smem.allocate_tensor(
        element_type=ab_dtype,
        layout=a_smem_layout_staged.outer,
        byte_alignment=128,
        swizzle=a_smem_layout_staged.inner,
    )
    # (MMA, MMA_N, MMA_K, STAGE)
    sB1 = smem.allocate_tensor(
        element_type=ab_dtype,
        layout=b_smem_layout_staged.outer,
        byte_alignment=128,
        swizzle=b_smem_layout_staged.inner,
    )
    # (MMA, MMA_N, MMA_K, STAGE)
    sB2 = smem.allocate_tensor(
        element_type=ab_dtype,
        layout=b_smem_layout_staged.outer,
        byte_alignment=128,
        swizzle=b_smem_layout_staged.inner,
    )
    # (MMA, MMA_M, MMA_K, STAGE)
    sSFA = smem.allocate_tensor(
        element_type=sf_dtype,
        layout=sfa_smem_layout_staged,
        byte_alignment=128,
    )
    # (MMA, MMA_N, MMA_K, STAGE)
    sSFB1 = smem.allocate_tensor(
        element_type=sf_dtype,
        layout=sfb_smem_layout_staged,
        byte_alignment=128,
    )
    # (MMA, MMA_N, MMA_K, STAGE)
    sSFB2 = smem.allocate_tensor(
        element_type=sf_dtype,
        layout=sfb_smem_layout_staged,
        byte_alignment=128,
    )

    #
    # Initialize mainloop ab_pipeline, acc_pipeline and their states
    #
    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=ab_pipeline_producer_group,
        consumer_group=pipeline.CooperativeGroup(
            pipeline.Agent.Thread,
            threads_per_cta,
        ),
    ).make_participants()

    #
    # Local_tile partition global tensors
    #
    # (bM, bK, RestM, RestK, RestL)
    gA_mkl = cute.local_tile(
        mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
    )
    # (bN, bK, RestN, RestK, RestL)
    gB_nkl1 = cute.local_tile(
        mB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )
    # (bN, bK, RestN, RestK, RestL)
    gB_nkl2 = cute.local_tile(
        mB_nkl2, 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_nkl1 = cute.local_tile(
        mSFB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )
    # (bN, bK, RestN, RestK, RestL)
    gSFB_nkl2 = cute.local_tile(
        mSFB_nkl2, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
    )
    # (bM, bN, RestM, RestN, RestL)
    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])

    #
    # Partition global tensor for TiledMMA_A/B/SFA/SFB/C
    #
    # (MMA, MMA_M, MMA_K, RestK)
    thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
    # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
    tCgA = thr_mma.partition_A(gA_mkl)
    # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
    tCgB1 = thr_mma.partition_B(gB_nkl1)
    # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
    tCgB2 = thr_mma.partition_B(gB_nkl2)
    # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
    tCgSFA = thr_mma.partition_A(gSFA_mkl)
    # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
    tCgSFB1 = thr_mma.partition_B(gSFB_nkl1)
    # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
    tCgSFB2 = thr_mma.partition_B(gSFB_nkl2)
    # (MMA, MMA_M, MMA_N, RestM, RestN, RestL)
    tCgC = thr_mma.partition_C(gC_mnl)

    #
    # Partition global/shared tensor for TMA load A/B/SFA/SFB
    #
    # TMA Partition_S/D for A
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestM, RestK, RestL)
    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),
    )
    # TMA Partition_S/D for B1
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestN, RestK, RestL)
    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),
    )
    # TMA Partition_S/D for B2
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestN, RestK, RestL)
    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),
    )
    #  TMA Partition_S/D for SFA
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestM, RestK, RestL)
    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)
    # TMA Partition_S/D for SFB1
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestN, RestK, RestL)
    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 Partition_S/D for SFB2
    # ((atom_v, rest_v), STAGE)
    # ((atom_v, rest_v), RestN, RestK, RestL)
    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)

    #
    # Partition shared/tensor memory tensor for TiledMMA_A/B/C
    #
    # (MMA, MMA_M, MMA_K, STAGE)
    tCrA = tiled_mma.make_fragment_A(sA)
    # (MMA, MMA_N, MMA_K, STAGE)
    tCrB1 = tiled_mma.make_fragment_B(sB1)
    # (MMA, MMA_N, MMA_K, STAGE)
    tCrB2 = tiled_mma.make_fragment_B(sB2)
    # (MMA, MMA_M, MMA_N)
    acc_shape = tiled_mma.partition_shape_C(mma_tiler_mnk[:2])
    # (MMA, MMA_M, MMA_N)
    tCtAcc_fake = tiled_mma.make_fragment_C(acc_shape)

    #
    # Alloc tensor memory buffer
    # Make ACC1 and ACC2 tmem tensor
    # ACC1 += A @ B1
    # ACC2 += A @ B2
    #
    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)

    #
    # Make SFA/SFB1/SFB2 tmem tensor
    #
    # SFA tmem layout: (MMA, MMA_M, MMA_K)
    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)),
    )
    # Get SFA tmem ptr
    sfa_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,
    )
    tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)

    # SFB1, SFB2 tmem layout: (MMA, MMA_N, MMA_K)
    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)),
    )
    # Get SFB1 tmem ptr
    sfb_tmem_ptr1 = 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(tCtSFA),
        dtype=sf_dtype,
    )
    tCtSFB1 = cute.make_tensor(sfb_tmem_ptr1, tCtSFB_layout)
    # Get SFB2 tmem ptr
    sfb_tmem_ptr2 = 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(tCtSFA)
        + tcgen05.find_tmem_tensor_col_offset(tCtSFB1),
        dtype=sf_dtype,
    )
    tCtSFB2 = cute.make_tensor(sfb_tmem_ptr2, tCtSFB_layout)

    #
    # Partition for S2T copy of SFA/SFB1/SFB2
    #
    # Make S2T CopyAtom
    copy_atom_s2t = cute.make_copy_atom(
        tcgen05.Cp4x32x128bOp(tcgen05.CtaGroup.ONE),
        sf_dtype,
    )
    # (MMA, MMA_MN, MMA_K, STAGE)
    tCsSFA_compact = cute.filter_zeros(sSFA)
    # (MMA, MMA_MN, MMA_K)
    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)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFA_compact_s2t_ = thr_copy_s2t_sfa.partition_S(tCsSFA_compact)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFA_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfa, tCsSFA_compact_s2t_
    )
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
    tCtSFA_compact_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA_compact)

    # (MMA, MMA_MN, MMA_K, STAGE)
    tCsSFB1_compact = cute.filter_zeros(sSFB1)
    # (MMA, MMA_MN, MMA_K)
    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)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFB1_compact_s2t_ = thr_copy_s2t_sfb.partition_S(tCsSFB1_compact)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFB1_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfb, tCsSFB1_compact_s2t_
    )
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
    tCtSFB1_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB1_compact)

    # SFB2 S2T copy and partition
    # (MMA, MMA_MN, MMA_K, STAGE)
    tCsSFB2_compact = cute.filter_zeros(sSFB2)
    # (MMA, MMA_MN, MMA_K)
    tCtSFB2_compact = cute.filter_zeros(tCtSFB2)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFB2_compact_s2t_ = thr_copy_s2t_sfb.partition_S(tCsSFB2_compact)
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
    tCsSFB2_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t_sfb, tCsSFB2_compact_s2t_
    )
    # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
    tCtSFB2_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB2_compact)

    #
    # Slice to per mma tile index
    #
    # ((atom_v, rest_v), RestK)
    tAgA = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
    # ((atom_v, rest_v), RestK)
    tBgB1 = tBgB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
    # ((atom_v, rest_v), RestK)
    tBgB2 = tBgB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
    # ((atom_v, rest_v), RestK)
    tAgSFA = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
    # ((atom_v, rest_v), RestK)
    tBgSFB1 = tBgSFB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
    # ((atom_v, rest_v), RestK)
    tBgSFB2 = tBgSFB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]

    #
    # Execute Data copy and Math computation in the k_tile loop
    #
    if warp_idx == 0:
        # Wait for accumulator buffer empty
        acc_empty = acc_producer.acquire_and_advance()
        # Set ACCUMULATE field to False for the first k_tile iteration
        tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
        # Execute k_tile loop
        for k_tile in range(k_tile_cnt):
            # Wait for AB buffer empty
            ab_empty = ab_producer.acquire_and_advance()

            #  TMA load A/B1/B2/SFA/SFB1/SFB2 to shared memory
            cute.copy(
                tma_atom_a,
                tAgA[(None, ab_empty.count)],
                tAsA[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )
            cute.copy(
                tma_atom_b1,
                tBgB1[(None, ab_empty.count)],
                tBsB1[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )
            cute.copy(
                tma_atom_b2,
                tBgB2[(None, ab_empty.count)],
                tBsB2[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )
            cute.copy(
                tma_atom_sfa,
                tAgSFA[(None, ab_empty.count)],
                tAsSFA[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )
            cute.copy(
                tma_atom_sfb1,
                tBgSFB1[(None, ab_empty.count)],
                tBsSFB1[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )
            cute.copy(
                tma_atom_sfb2,
                tBgSFB2[(None, ab_empty.count)],
                tBsSFB2[(None, ab_empty.index)],
                tma_bar_ptr=ab_empty.barrier,
            )

            # Wait for AB buffer full
            ab_full = ab_consumer.wait_and_advance()

            #  Copy SFA/SFB1/SFB2 to tmem
            s2t_stage_coord = (None, None, None, None, ab_full.index)
            tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
            tCsSFB1_compact_s2t_staged = tCsSFB1_compact_s2t[s2t_stage_coord]
            tCsSFB2_compact_s2t_staged = tCsSFB2_compact_s2t[s2t_stage_coord]
            cute.copy(
                tiled_copy_s2t_sfa,
                tCsSFA_compact_s2t_staged,
                tCtSFA_compact_s2t,
            )
            cute.copy(
                tiled_copy_s2t_sfb,
                tCsSFB1_compact_s2t_staged,
                tCtSFB1_compact_s2t,
            )
            cute.copy(
                tiled_copy_s2t_sfb,
                tCsSFB2_compact_s2t_staged,
                tCtSFB2_compact_s2t,
            )

            # tCtAcc1 += tCrA * tCrSFA * tCrB1 * tCrSFB1
            # tCtAcc2 += tCrA * tCrSFA * tCrB2 * tCrSFB2
            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,
                )

                # Set SFA/SFB tensor to tiled_mma
                sf_kblock_coord = (None, None, kblock_idx)
                tiled_mma.set(
                    tcgen05.Field.SFA,
                    tCtSFA[sf_kblock_coord].iterator,
                )
                tiled_mma.set(
                    tcgen05.Field.SFB,
                    tCtSFB1[sf_kblock_coord].iterator,
                )
                cute.gemm(
                    tiled_mma,
                    tCtAcc1,
                    tCrA[kblock_coord],
                    tCrB1[kblock_coord],
                    tCtAcc1,
                )

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

                # Enable accumulate on tCtAcc1/tCtAcc2 after first kblock
                tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

            # Async arrive AB buffer empty
            ab_full.release()
        acc_empty.commit()

    #
    # Epilogue
    # Partition for 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)
    # (T2R_M, T2R_N, EPI_M, EPI_M)
    tTR_tAcc1 = thr_copy_t2r.partition_S(tCtAcc1)
    # (T2R_M, T2R_N, EPI_M, EPI_M)
    tTR_tAcc2 = thr_copy_t2r.partition_S(tCtAcc2)
    # (T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
    tTR_gC = thr_copy_t2r.partition_D(tCgC)
    # (T2R_M, T2R_N, EPI_M, EPI_N)
    tTR_rAcc1 = cute.make_rmem_tensor(
        tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32
    )
    # (T2R_M, T2R_N, EPI_M, EPI_N)
    tTR_rAcc2 = cute.make_rmem_tensor(
        tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32
    )
    # (T2R_M, T2R_N, EPI_M, EPI_N)
    tTR_rC = cute.make_rmem_tensor(
        tTR_gC[None, None, None, None, 0, 0, 0].shape, c_dtype
    )
    # STG Atom
    simt_atom = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), c_dtype)
    tTR_gC = tTR_gC[(None, None, None, None, *mma_tile_coord_mnl)]

    # Wait for accumulator buffer full
    acc_full = acc_consumer.wait_and_advance()

    # Copy accumulator to register
    cute.copy(tiled_copy_t2r, tTR_tAcc1, tTR_rAcc1)
    cute.copy(tiled_copy_t2r, tTR_tAcc2, tTR_rAcc2)

    # Silu activation on acc1 and multiply with acc2
    acc_vec1 = epilogue_op(tTR_rAcc1.load())
    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)

    acc_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))),
):
    """
    Host-side JIT-compiled function that prepares tensors and launches the GPU kernel.
    
    This function is compiled just-in-time by CUTLASS and serves as the bridge between
    the Python host code and the GPU device kernel. It performs the following setup:
    
    1. Creates CuTe tensor descriptors from raw pointers with appropriate layouts:
       - A tensor: (M, K, L) with K-major layout (contiguous along K dimension)
       - B1/B2 tensors: (N, K, L) with K-major layout
       - C tensor: (M, N, L) with N-major layout (contiguous along N dimension)
       - Scale factor tensors: Special block-scaled layout for FP4 quantization
    
    2. Configures the MMA (Matrix Multiply-Accumulate) operation:
       - Uses tcgen05.MmaMXF4NVF4Op for Blackwell's NVFP4 tensor cores
       - Sets tile shape (128x128x64) matching the hardware MMA instruction
    
    3. Computes shared memory layouts with swizzling for bank-conflict-free access
    
    4. Sets up TMA (Tensor Memory Accelerator) descriptors for efficient bulk transfers:
       - TMA enables asynchronous global-to-shared memory copies
       - Each TMA atom defines how to transfer one tile of data
    
    5. Calculates grid dimensions and launches the device kernel
    
    Args:
        a_ptr: CuTe pointer to matrix A data in global memory (FP4 format)
        b1_ptr: CuTe pointer to matrix B1 data in global memory (FP4 format)
        b2_ptr: CuTe pointer to matrix B2 data in global memory (FP4 format)
        sfa_ptr: CuTe pointer to scale factors for A (FP8 format, block-permuted layout)
        sfb1_ptr: CuTe pointer to scale factors for B1 (FP8 format, block-permuted layout)
        sfb2_ptr: CuTe pointer to scale factors for B2 (FP8 format, block-permuted layout)
        c_ptr: CuTe pointer to output matrix C in global memory (FP16 format)
        problem_size: Tuple of (M, N, K, L) dimensions where:
                      M = number of rows in A and C
                      N = number of columns in B1/B2 and C
                      K = shared/reduction dimension
                      L = batch size
        epilogue_op: Element-wise operation for activation (default: SiLU)
    
    Note:
        The scale factor pointers expect data in a special permuted layout that matches
        the hardware's block-scaled MMA requirements. This permutation is done in the
        task.py preprocessing step.
    """
    m, n, k, l = problem_size

    # Setup attributes that depend on gemm inputs
    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))
    )
    # Setup sfa/sfb tensor by filling A/B tensor to scale factor atom layout
    # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL)
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
        a_tensor.shape, sf_vec_size
    )
    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

    # ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL)
    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,),
    )

    # Compute A/B/SFA/SFB/C shared memory layout
    a_smem_layout_staged = sm100_utils.make_smem_layout_a(
        tiled_mma,
        mma_tiler_mnk,
        ab_dtype,
        num_ab_stage,
    )
    # B1 and B2 have the same size thus share the same smem layout
    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,
    )
    # SFB1 and SFB2 have the same size thus share the same smem layout
    sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
        tiled_mma,
        mma_tiler_mnk,
        sf_vec_size,
        num_ab_stage,
    )
    atom_thr_size = cute.size(tiled_mma.thr_id.shape)

    # Setup TMA for A
    a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, None, 0))
    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,
        cluster_layout_vmnk .shape,
    )
    # Setup TMA for B1
    b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, None, 0))
    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,
    )
    # Setup TMA for B2
    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,
    )
    # Setup TMA for SFA
    sfa_smem_layout = cute.slice_(
        sfa_smem_layout_staged , (None, None, None, 0)
    )
    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,
        cluster_layout_vmnk .shape,
        internal_type=cutlass.Int16,
    )
    # Setup TMA for SFB1
    sfb_smem_layout = cute.slice_(
        sfb_smem_layout_staged , (None, None, None, 0)
    )
    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,
    )
    # Setup TMA for SFB2
    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,
    )

    # Compute TMA load bytes
    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 * 2 + sfa_copy_size + sfb_copy_size * 2
    ) * atom_thr_size

    # Compute grid 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],
    )

    # Launch the kernel.
    kernel(
        # MMA (Matrix Multiply-Accumulate) configuration
        tiled_mma,                  # Tiled MMA object defining NVFP4 GEMM compute pattern
        
        # TMA (Tensor Memory Accelerator) atoms and tensors for shared input matrix A
        tma_atom_a,                 # TMA copy atom defining how to load A from global memory
        tma_tensor_a,               # Tensor descriptor for A matrix (m, k, l) - shared by both GEMMs
        
        # TMA atoms and tensors for first B matrix (B1)
        tma_atom_b1,                # TMA copy atom defining how to load B1 from global memory
        tma_tensor_b1,              # Tensor descriptor for B1 matrix (n, k, l) - first GEMM
        
        # TMA atoms and tensors for second B matrix (B2)
        tma_atom_b2,                # TMA copy atom defining how to load B2 from global memory
        tma_tensor_b2,              # Tensor descriptor for B2 matrix (n, k, l) - second GEMM
        
        # TMA atoms and tensors for scale factor A (shared)
        tma_atom_sfa,               # TMA copy atom for loading scale factors for A
        tma_tensor_sfa,             # Tensor descriptor for SFA (block scale factors for A) - shared
        
        # TMA atoms and tensors for scale factor B1
        tma_atom_sfb1,              # TMA copy atom for loading scale factors for B1
        tma_tensor_sfb1,            # Tensor descriptor for SFB1 (block scale factors for B1)
        
        # TMA atoms and tensors for scale factor B2
        tma_atom_sfb2,              # TMA copy atom for loading scale factors for B2
        tma_tensor_sfb2,            # Tensor descriptor for SFB2 (block scale factors for B2)
        
        # Output tensor C (stores both C1 and C2 results)
        c_tensor,                   # Output tensor where both GEMM results will be stored (m, n, l)
        
        # Shared memory layouts with staging for pipelined execution
        a_smem_layout_staged,       # Staged shared memory layout for A (includes stage dimension)
        b_smem_layout_staged,       # Staged shared memory layout for B1/B2 (includes stage dimension)
        sfa_smem_layout_staged,     # Staged shared memory layout for SFA (includes stage dimension)
        sfb_smem_layout_staged,     # Staged shared memory layout for SFB1/SFB2 (includes stage dimension)
        
        # Pipeline synchronization parameter
        num_tma_load_bytes,         # Total bytes to load per TMA transaction (for barrier setup)
        
        # Epilogue operation
        epilogue_op,                # Epilogue operation to apply to output (e.g., element-wise ops)
    ).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():
    """
    Compile the dual GEMM kernel once and cache it for repeated execution.
    
    This function implements a singleton pattern for kernel compilation to avoid
    the overhead of JIT compilation during performance benchmarking. The first call
    triggers full compilation, while subsequent calls return the cached kernel.
    
    The compilation process:
    1. Creates dummy CuTe pointers with the correct data types (FP4, FP8, FP16)
    2. Invokes cute.compile() which triggers CUTLASS's JIT compilation pipeline
    3. Generates PTX/SASS code optimized for the target GPU (Blackwell)
    4. Caches the compiled function for reuse
    
    Why this matters for benchmarking:
    - JIT compilation can take several seconds
    - Benchmark frameworks typically call custom_kernel() multiple times
    - Without caching, each call would re-compile, skewing timing results
    
    Returns:
        Callable: The compiled kernel function that can be invoked with actual
                  tensor pointers and problem dimensions
    
    Usage:
        # Pre-compile before timing
        compiled_func = compile_kernel()
        
        # Time only the kernel execution
        start = time.time()
        compiled_func(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, dims)
        torch.cuda.synchronize()
        elapsed = time.time() - start
    """
    global _compiled_kernel_cache
    
    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache
    

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    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
    )

    # Compile the kernel
    _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:
    """
    Execute the block-scaled dual GEMM kernel with SiLU-gated activation.
    
    Computes: C = silu(A @ B1) * (A @ B2)
    
    This is the main entry point called by the evaluation/benchmarking framework.
    It handles the conversion from PyTorch tensors to CuTe pointers and invokes
    the pre-compiled CUDA kernel.
    
    Data flow:
    1. Unpack input tensors from the data tuple
    2. Retrieve or compile the cached kernel
    3. Extract problem dimensions (M, N, K, L)
    4. Create CuTe pointers from PyTorch tensor data pointers
    5. Execute the compiled kernel with the actual data
    6. Return the output tensor (modified in-place)
    
    Memory layout notes:
    - A, B1, B2 are stored in K-major order (contiguous along K)
    - C is stored in N-major order (contiguous along N)
    - Scale factors use a special 6D permuted layout for hardware compatibility:
      [32, 4, rest_dim, 4, rest_k, batch] which matches the block-scaled MMA's
      expected access pattern
    
    Args:
        data: Tuple containing all input/output tensors:
            a: [M, K//2, L] - Input matrix A in packed FP4 (float4e2m1fn_x2)
                              K is halved because 2 FP4 values are packed per element
            b1: [N, K//2, L] - Input matrix B1 in packed FP4
            b2: [N, K//2, L] - Input matrix B2 in packed FP4
            sfa_cpu: [M//16, K//16, L] - Scale factors for A (unpermuted, unused here)
            sfb1_cpu: [N//16, K//16, L] - Scale factors for B1 (unpermuted, unused here)
            sfb2_cpu: [N//16, K//16, L] - Scale factors for B2 (unpermuted, unused here)
            sfa_permuted: [32, 4, M//128, 4, K//256, L] - Permuted scale factors for A
            sfb1_permuted: [32, 4, N//128, 4, K//256, L] - Permuted scale factors for B1
            sfb2_permuted: [32, 4, N//128, 4, K//256, L] - Permuted scale factors for B2
            c: [M, N, L] - Output matrix in FP16, modified in-place
    
    Returns:
        output_t: The output tensor C containing silu(A @ B1) * (A @ B2)
    
    Note:
        The kernel modifies C in-place. The same tensor passed in is returned,
        now containing the computed result.
    """
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
    
    # Ensure kernel is compiled (will use cached version if available)
    # To avoid the compilation overhead, we compile the kernel once and cache it.
    compiled_func = compile_kernel()

    # Get dimensions from MxKxL layout
    _, k, _ = a.shape
    m, n, l = c.shape
    # Torch use e2m1_x2 data type, thus k is halved
    k = k * 2 

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    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
    )

    # Execute the compiled kernel
    compiled_func(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l))

    return c
scrolls · 1103 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 356803.

⋯ 37 unchanged lines
# Helper function for ceiling division
def ceil_div(a, b):
+ """
+ Compute the ceiling of a / b using integer arithmetic.
+
+ This is used throughout the kernel to calculate grid dimensions and tile counts
+ when the problem size doesn't evenly divide by the tile size.
+
+ Args:
+ a: Numerator (e.g., problem dimension M, N, or K)
+ b: Denominator (e.g., tile size)
+
+ Returns:
+ ceil(a / b) as an integer
+
+ Example:
+ ceil_div(1000, 128) = 8 # Need 8 tiles of size 128 to cover 1000 elements
+ """
return (a + b - 1) // b
⋯ 23 unchanged lines
* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
"""
- GPU device kernel performing the batched GEMM computation.
+ GPU device kernel performing dual block-scaled NVFP4 GEMM with SiLU-gated fusion.
+
+ Computes: C = silu(A @ B1) * (A @ B2)
+
+ This kernel executes the following operations:
+ 1. Loads FP4 matrices A, B1, B2 and their FP8 block scale factors from global memory
+ to shared memory using TMA (Tensor Memory Accelerator)
+ 2. Copies scale factors from shared memory to tensor memory (TMEM)
+ 3. Performs two block-scaled matrix multiplications using Blackwell's tcgen05 MMA:
+ - ACC1 = A @ B1 (with block scaling from SFA and SFB1)
+ - ACC2 = A @ B2 (with block scaling from SFA and SFB2)
+ 4. Applies SiLU activation to ACC1 and multiplies element-wise with ACC2
+ 5. Stores the FP16 result to global memory
+
+ Key optimization: Matrix A and its scale factors (SFA) are loaded once but reused
+ for both GEMMs, reducing memory bandwidth by ~33%.
+
+ Execution model:
+ - Warp 0 handles all TMA loads and MMA operations (producer)
+ - All warps participate in the epilogue (copying results to global memory)
+ - Pipeline barriers synchronize TMA loads with MMA consumption
+
+ Args:
+ tiled_mma: TiledMMA object defining the MMA instruction configuration
+ (tile shape, data types, tensor core operation)
+ tma_atom_a: TMA copy atom for loading matrix A from global to shared memory
+ mA_mkl: Global memory tensor for A with shape (M, K, L) in FP4 format
+ tma_atom_b1: TMA copy atom for loading matrix B1
+ mB_nkl1: Global memory tensor for B1 with shape (N, K, L) in FP4 format
+ tma_atom_b2: TMA copy atom for loading matrix B2
+ mB_nkl2: Global memory tensor for B2 with shape (N, K, L) in FP4 format
+ tma_atom_sfa: TMA copy atom for loading scale factors for A
+ mSFA_mkl: Global memory tensor for SFA (block scale factors for A)
+ tma_atom_sfb1: TMA copy atom for loading scale factors for B1
+ mSFB_nkl1: Global memory tensor for SFB1 (block scale factors for B1)
+ tma_atom_sfb2: TMA copy atom for loading scale factors for B2
+ mSFB_nkl2: Global memory tensor for SFB2 (block scale factors for B2)
+ mC_mnl: Global memory tensor for output C with shape (M, N, L) in FP16
+ a_smem_layout_staged: Shared memory layout for A with staging dimension
+ b_smem_layout_staged: Shared memory layout for B1/B2 with staging
+ sfa_smem_layout_staged: Shared memory layout for SFA with staging
+ sfb_smem_layout_staged: Shared memory layout for SFB1/SFB2 with staging
+ num_tma_load_bytes: Total bytes per TMA transaction (for barrier setup)
+ epilogue_op: Element-wise operation applied to ACC1 (default: SiLU activation)
"""
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
⋯ 552 unchanged lines
* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
"""
- Host-side JIT function to prepare tensors and launch GPU kernel.
+ Host-side JIT-compiled function that prepares tensors and launches the GPU kernel.
+
+ This function is compiled just-in-time by CUTLASS and serves as the bridge between
+ the Python host code and the GPU device kernel. It performs the following setup:
+
+ 1. Creates CuTe tensor descriptors from raw pointers with appropriate layouts:
+ - A tensor: (M, K, L) with K-major layout (contiguous along K dimension)
+ - B1/B2 tensors: (N, K, L) with K-major layout
+ - C tensor: (M, N, L) with N-major layout (contiguous along N dimension)
+ - Scale factor tensors: Special block-scaled layout for FP4 quantization
+
+ 2. Configures the MMA (Matrix Multiply-Accumulate) operation:
+ - Uses tcgen05.MmaMXF4NVF4Op for Blackwell's NVFP4 tensor cores
+ - Sets tile shape (128x128x64) matching the hardware MMA instruction
+
+ 3. Computes shared memory layouts with swizzling for bank-conflict-free access
+
+ 4. Sets up TMA (Tensor Memory Accelerator) descriptors for efficient bulk transfers:
+ - TMA enables asynchronous global-to-shared memory copies
+ - Each TMA atom defines how to transfer one tile of data
+
+ 5. Calculates grid dimensions and launches the device kernel
+
+ Args:
+ a_ptr: CuTe pointer to matrix A data in global memory (FP4 format)
+ b1_ptr: CuTe pointer to matrix B1 data in global memory (FP4 format)
+ b2_ptr: CuTe pointer to matrix B2 data in global memory (FP4 format)
+ sfa_ptr: CuTe pointer to scale factors for A (FP8 format, block-permuted layout)
+ sfb1_ptr: CuTe pointer to scale factors for B1 (FP8 format, block-permuted layout)
+ sfb2_ptr: CuTe pointer to scale factors for B2 (FP8 format, block-permuted layout)
+ c_ptr: CuTe pointer to output matrix C in global memory (FP16 format)
+ problem_size: Tuple of (M, N, K, L) dimensions where:
+ M = number of rows in A and C
+ N = number of columns in B1/B2 and C
+ K = shared/reduction dimension
+ L = batch size
+ epilogue_op: Element-wise operation for activation (default: SiLU)
+
+ Note:
+ The scale factor pointers expect data in a special permuted layout that matches
+ the hardware's block-scaled MMA requirements. This permutation is done in the
+ task.py preprocessing step.
"""
m, n, k, l = problem_size
⋯ 213 unchanged lines
# Global cache for compiled kernel
_compiled_kernel_cache = None
- # This function is used to compile the kernel once and cache it and then allow users to
- # run the kernel multiple times to get more accurate timing results.
+
+
def compile_kernel():
"""
- Compile the kernel once and cache it.
- This should be called before any timing measurements.
+ Compile the dual GEMM kernel once and cache it for repeated execution.
+ This function implements a singleton pattern for kernel compilation to avoid
+ the overhead of JIT compilation during performance benchmarking. The first call
+ triggers full compilation, while subsequent calls return the cached kernel.
+
+ The compilation process:
+ 1. Creates dummy CuTe pointers with the correct data types (FP4, FP8, FP16)
+ 2. Invokes cute.compile() which triggers CUTLASS's JIT compilation pipeline
+ 3. Generates PTX/SASS code optimized for the target GPU (Blackwell)
+ 4. Caches the compiled function for reuse
+
+ Why this matters for benchmarking:
+ - JIT compilation can take several seconds
+ - Benchmark frameworks typically call custom_kernel() multiple times
+ - Without caching, each call would re-compile, skewing timing results
+
Returns:
- The compiled kernel function
+ Callable: The compiled kernel function that can be invoked with actual
+ tensor pointers and problem dimensions
+
+ Usage:
+ # Pre-compile before timing
+ compiled_func = compile_kernel()
+
+ # Time only the kernel execution
+ start = time.time()
+ compiled_func(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, dims)
+ torch.cuda.synchronize()
+ elapsed = time.time() - start
"""
global _compiled_kernel_cache
⋯ 32 unchanged lines
def custom_kernel(data: input_t) -> output_t:
"""
- Execute the block-scaled dual GEMM kernel with silu activation,
- C = silu(A @ B1) * (A @ B2).
+ Execute the block-scaled dual GEMM kernel with SiLU-gated activation.
- This is the main entry point called by the evaluation framework.
- It converts PyTorch tensors to CuTe tensors, launches the kernel,
- and returns the result.
+ Computes: C = silu(A @ B1) * (A @ B2)
+ This is the main entry point called by the evaluation/benchmarking framework.
+ It handles the conversion from PyTorch tensors to CuTe pointers and invokes
+ the pre-compiled CUDA kernel.
+
+ Data flow:
+ 1. Unpack input tensors from the data tuple
+ 2. Retrieve or compile the cached kernel
+ 3. Extract problem dimensions (M, N, K, L)
+ 4. Create CuTe pointers from PyTorch tensor data pointers
+ 5. Execute the compiled kernel with the actual data
+ 6. Return the output tensor (modified in-place)
+
+ Memory layout notes:
+ - A, B1, B2 are stored in K-major order (contiguous along K)
+ - C is stored in N-major order (contiguous along N)
+ - Scale factors use a special 6D permuted layout for hardware compatibility:
+ [32, 4, rest_dim, 4, rest_k, batch] which matches the block-scaled MMA's
+ expected access pattern
+
Args:
- data: Tuple of (a, b1, b2, sfa_cpu, sfb1_cpu, sfb2_cpu, c) PyTorch tensors
- a: [m, k, l] - Input matrix in float4e2m1fn
- b1: [n, k, l] - Input matrix in float4e2m1fn
- b2: [n, k, l] - Input matrix in float4e2m1fn
- sfa_cpu: [m, k, l] - Scale factors in float8_e4m3fn, used by reference implementation
- sfb1_cpu: [n, k, l] - Scale factors in float8_e4m3fn, used by reference implementation
- sfb2_cpu: [n, k, l] - Scale factors in float8_e4m3fn, used by reference implementation
- sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn
- sfb1_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
- sfb2_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
- c: [m, n, l] - Output vector in float16
+ data: Tuple containing all input/output tensors:
+ a: [M, K//2, L] - Input matrix A in packed FP4 (float4e2m1fn_x2)
+ K is halved because 2 FP4 values are packed per element
+ b1: [N, K//2, L] - Input matrix B1 in packed FP4
+ b2: [N, K//2, L] - Input matrix B2 in packed FP4
+ sfa_cpu: [M//16, K//16, L] - Scale factors for A (unpermuted, unused here)
+ sfb1_cpu: [N//16, K//16, L] - Scale factors for B1 (unpermuted, unused here)
+ sfb2_cpu: [N//16, K//16, L] - Scale factors for B2 (unpermuted, unused here)
+ sfa_permuted: [32, 4, M//128, 4, K//256, L] - Permuted scale factors for A
+ sfb1_permuted: [32, 4, N//128, 4, K//256, L] - Permuted scale factors for B1
+ sfb2_permuted: [32, 4, N//128, 4, K//256, L] - Permuted scale factors for B2
+ c: [M, N, L] - Output matrix in FP16, modified in-place
Returns:
- Output tensor c with computed results
+ output_t: The output tensor C containing silu(A @ B1) * (A @ B2)
+
+ Note:
+ The kernel modifies C in-place. The same tensor passed in is returned,
+ now containing the computed result.
"""
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
scrolls · 238 diff lines total

Best evidence level for this revision: reported

JSON