Skip to content
KernelIndex
Search⌘K

submission 155546

rcmalli · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-155546?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 GEMMsuite of 3 cases
NVIDIA B200
12.7µs
#111 of 369
2025-12-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a71118bbe43d1770fe102e64d28278c898556ceb6447044eb7bd97b0b7cf58c0
license declaredunknown
license concludedunknown
authorsrcmalli
imported2026-08-15

Techniques

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

fp4Custom NVFP4 block-scaled GEMM kernel implementation using CuTe-DSL.
fused-epilogue- Warp specialization: TMA load (warp 5), MMA (warp 4), Epilogue (warps 0-3)
mbarrierself.epilog_sync_barrier = pipeline.NamedBarrier(
persistent-kernelThis implements a persistent warp-specialized kernel for B200 (SM100) targeting
shared-memoryself.smem_capacity = 228 * 1024 # 233472 bytes
tcgen05tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission.py1725 lines
"""
Custom NVFP4 block-scaled GEMM kernel implementation using CuTe-DSL.

This implements a persistent warp-specialized kernel for B200 (SM100) targeting
the NVFP4 block-scaled GEMM competition.

Key features:
- Warp specialization: TMA load (warp 5), MMA (warp 4), Epilogue (warps 0-3)
- Block-scaled FP4 with FP8 scale factors (sf_vec_size=16)
- TMA multicast for B matrix across cluster
- Software pipelining for A/B/SF data
"""

from dataclasses import dataclass
from typing import Optional, Tuple, Type, Union

import os

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

from task import input_t, output_t


@dataclass(frozen=True)
class SchedulerParams:
    tiles_m: int
    tiles_n: int
    k_tile_cnt: int

    @property
    def tiles_mn(self) -> int:
        return self.tiles_m * self.tiles_n

    def total_work_items(self, l: int) -> int:
        return self.tiles_mn * l

    @classmethod
    def from_shape(
        cls,
        m: int,
        n: int,
        k_tile_cnt: int,
        tile_m: int,
        tile_n: int,
    ):
        tiles_m = (m + tile_m - 1) // tile_m
        tiles_n = (n + tile_n - 1) // tile_n
        return cls(tiles_m, tiles_n, k_tile_cnt)


class Sm100BlockScaledGemmKernel:
    """
    Block-scaled FP4 GEMM kernel for SM100 (Blackwell).

    Configuration:
    - mma_tiler_mn: (128, 128) for M=128 cases
    - cluster_shape_mn: (1, 2) for B-matrix multicast
    - sf_vec_size: 16 (standard for NVFP4)
    """

    def __init__(
        self,
        sf_vec_size: int = 16,
        mma_tiler_mn: Tuple[int, int] = (128, 128),
        cluster_shape_mn: Tuple[int, int] = (1, 2),
        num_ab_stage: Optional[int] = None,
        num_c_stage: Optional[int] = None,
        ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN,
        sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN,
        c_dtype: Type[cutlass.Numeric] = cutlass.Float16,
        use_persistent_scheduler: bool = True,
    ):
        self.acc_dtype = cutlass.Float32
        self.sf_vec_size = sf_vec_size
        self.use_2cta_instrs = mma_tiler_mn[0] == 256
        self.cluster_shape_mn = cluster_shape_mn
        self.mma_tiler = (*mma_tiler_mn, 1)  # K dimension set later
        self.use_persistent_scheduler = use_persistent_scheduler
        self.num_ab_stage_override = num_ab_stage
        self.num_c_stage_override = num_c_stage

        # Store data types as instance variables (needed for JIT compilation)
        self.a_dtype = ab_dtype
        self.b_dtype = ab_dtype
        self.sf_dtype = sf_dtype
        self.c_dtype = c_dtype

        self.cta_group = (
            tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
        )

        self.occupancy = 1

        # Warp specialization IDs
        self.epilog_warp_id = (0, 1, 2, 3)
        self.mma_warp_id = 4
        self.tma_warp_id = 5
        self.threads_per_cta = 32 * len(
            (self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
        )

        # Named barriers
        self.epilog_sync_barrier = pipeline.NamedBarrier(
            barrier_id=1,
            num_threads=32 * len(self.epilog_warp_id),
        )
        self.tmem_alloc_barrier = pipeline.NamedBarrier(
            barrier_id=2,
            num_threads=self.threads_per_cta,
        )
        # Keep all warps/roles in lockstep across work items (ported from dev_v6 Strm-K loop)
        self.work_loop_barrier = pipeline.NamedBarrier(
            barrier_id=3,
            num_threads=self.threads_per_cta,
        )

        # SM100 has 228KB shared memory - hardware constant
        self.smem_capacity = 228 * 1024  # 233472 bytes
        SM100_TMEM_CAPACITY_COLUMNS = 512
        self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS

    def _setup_attributes(self):
        """Configure kernel attributes based on input tensor properties."""
        # MMA instruction shape
        self.mma_inst_shape_mn = (self.mma_tiler[0], self.mma_tiler[1])
        self.mma_inst_shape_mn_sfb = (
            self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
            cute.round_up(self.mma_inst_shape_mn[1], 128),
        )

        # Create tiled MMA for block-scaled operation
        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype,
            self.a_major_mode,
            self.b_major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            self.cta_group,
            self.mma_inst_shape_mn,
        )

        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype,
            self.a_major_mode,
            self.b_major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            tcgen05.CtaGroup.ONE,
            self.mma_inst_shape_mn_sfb,
        )

        # Compute MMA/cluster/tile shapes
        mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
        mma_inst_tile_k = 4
        self.mma_tiler = (
            self.mma_inst_shape_mn[0],
            self.mma_inst_shape_mn[1],
            mma_inst_shape_k * mma_inst_tile_k,
        )
        self.mma_tiler_sfb = (
            self.mma_inst_shape_mn_sfb[0],
            self.mma_inst_shape_mn_sfb[1],
            mma_inst_shape_k * mma_inst_tile_k,
        )
        self.cta_tile_shape_mnk = (
            self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
            self.mma_tiler[1],
            self.mma_tiler[2],
        )
        self.cta_tile_shape_mnk_sfb = (
            self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape),
            self.mma_tiler_sfb[1],
            self.mma_tiler_sfb[2],
        )

        # Cluster layout
        self.cluster_layout_vmnk = cute.tiled_divide(
            cute.make_layout((*self.cluster_shape_mn, 1)),
            (tiled_mma.thr_id.shape,),
        )
        self.cluster_layout_sfb_vmnk = cute.tiled_divide(
            cute.make_layout((*self.cluster_shape_mn, 1)),
            (tiled_mma_sfb.thr_id.shape,),
        )

        # Multicast configuration
        self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2])
        self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1])
        self.num_mcast_ctas_sfb = cute.size(self.cluster_layout_sfb_vmnk.shape[1])
        self.is_a_mcast = self.num_mcast_ctas_a > 1
        self.is_b_mcast = self.num_mcast_ctas_b > 1
        self.is_sfb_mcast = self.num_mcast_ctas_sfb > 1

        # Epilogue tile
        self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
            self.cta_tile_shape_mnk,
            self.use_2cta_instrs,
            self.c_layout,
            self.c_dtype,
        )
        self.epi_tile_n = cute.size(self.epi_tile[1])

        # Stage counts
        self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
            tiled_mma,
            self.mma_tiler,
            self.a_dtype,
            self.b_dtype,
            self.epi_tile,
            self.c_dtype,
            self.c_layout,
            self.sf_dtype,
            self.sf_vec_size,
            self.smem_capacity,
            self.occupancy,
        )
        if self.num_ab_stage_override is not None and self.num_ab_stage_override > 0:
            self.num_ab_stage = max(2, min(8, int(self.num_ab_stage_override)))
        if self.num_c_stage_override is not None and self.num_c_stage_override > 0:
            self.num_c_stage = max(1, min(4, int(self.num_c_stage_override)))

        # SMEM layouts
        self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(
            tiled_mma,
            self.mma_tiler,
            self.a_dtype,
            self.num_ab_stage,
        )
        self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(
            tiled_mma,
            self.mma_tiler,
            self.b_dtype,
            self.num_ab_stage,
        )
        self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
        )
        self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
        )
        self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
            self.c_dtype,
            self.c_layout,
            self.epi_tile,
            self.num_c_stage,
        )

        # Accumulator configuration
        self.overlapping_accum = self.num_acc_stage == 1

        # TMEM column counts
        sf_atom_mn = 32
        self.num_sfa_tmem_cols = (
            self.cta_tile_shape_mnk[0] // sf_atom_mn
        ) * mma_inst_tile_k
        self.num_sfb_tmem_cols = (
            self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
        ) * mma_inst_tile_k
        self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
        self.num_accumulator_tmem_cols = (
            self.cta_tile_shape_mnk[1] * self.num_acc_stage
            if not self.overlapping_accum
            else self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols
        )
        self.iter_acc_early_release_in_epilogue = (
            self.num_sf_tmem_cols // self.epi_tile_n
        )

    def _compute_stages(
        self,
        tiled_mma,
        mma_tiler,
        a_dtype,
        b_dtype,
        epi_tile,
        c_dtype,
        c_layout,
        sf_dtype,
        sf_vec_size,
        smem_capacity,
        occupancy,
    ):
        """Compute pipeline stage counts based on SMEM capacity."""
        # Compute SMEM footprints
        a_smem_layout = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler, a_dtype, 1)
        b_smem_layout = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler, b_dtype, 1)
        sfa_smem_layout = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma, mma_tiler, sf_vec_size, 1
        )
        sfb_smem_layout = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma, mma_tiler, sf_vec_size, 1
        )
        c_smem_layout = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)

        a_bytes = cute.size_in_bytes(a_dtype, a_smem_layout.outer)
        b_bytes = cute.size_in_bytes(b_dtype, b_smem_layout.outer)
        sfa_bytes = cute.size_in_bytes(sf_dtype, sfa_smem_layout)
        sfb_bytes = cute.size_in_bytes(sf_dtype, sfb_smem_layout)
        c_bytes = cute.size_in_bytes(c_dtype, c_smem_layout.outer)

        ab_sf_bytes_per_stage = a_bytes + b_bytes + sfa_bytes + sfb_bytes

        # Force num_acc_stage=1 to enable overlapping_accum (accumulator ping-pong)
        # This allows epilogue to overlap with next tile's MMA via early release
        num_acc_stage = 1

        # Compute C stages (2 for double buffering)
        num_c_stage = 2

        # Fixed overhead for barriers, etc.
        barrier_bytes = 1024

        # Available SMEM for A/B/SF stages
        available = smem_capacity // occupancy - c_bytes * num_c_stage - barrier_bytes
        num_ab_stage = max(2, min(8, available // ab_sf_bytes_per_stage))

        return num_acc_stage, num_ab_stage, num_c_stage

    def _compute_grid(
        self, c_tensor, cta_tile_shape, cluster_shape_mn, max_active_clusters
    ):
        """Compute grid dimensions for kernel launch.

        Uses CuTe's zipped_divide to compute the number of CTAs needed,
        then uses StaticPersistentTileScheduler.get_grid_shape for the grid.
        """
        # Use zipped_divide to get the number of CTAs per dimension
        c_shape = cute.slice_(cta_tile_shape, (None, None, 0))
        gc = cute.zipped_divide(c_tensor, tiler=c_shape)
        num_ctas_mnl = gc[(0, (None, None, None))].shape

        # PersistentTileSchedulerParams expects cluster_shape_mnk with K=1
        cluster_shape_mnk = (*cluster_shape_mn, 1)

        tile_sched_params = utils.PersistentTileSchedulerParams(
            num_ctas_mnl, cluster_shape_mnk
        )

        # Use StaticPersistentTileScheduler to compute the grid shape
        grid = utils.StaticPersistentTileScheduler.get_grid_shape(
            tile_sched_params, max_active_clusters
        )

        return tile_sched_params, grid

    @cute.jit
    def __call__(
        self,
        a_ptr: cute.Pointer,
        b_ptr: cute.Pointer,
        sfa_ptr: cute.Pointer,
        sfb_ptr: cute.Pointer,
        c_ptr: cute.Pointer,
        shape_mnkl: Tuple[int, int, int, int],
        max_active_clusters: cutlass.Constexpr = 132,
        epilogue_op: cutlass.Constexpr = lambda x: x,
    ):
        """Execute the block-scaled GEMM operation.

        Args:
            a_ptr: Pointer to A matrix data (M x K x L, K-major, FP4)
            b_ptr: Pointer to B matrix data (N x K x L, K-major, FP4)
            sfa_ptr: Pointer to scale factors for A (permuted layout, FP8)
            sfb_ptr: Pointer to scale factors for B (permuted layout, FP8)
            c_ptr: Pointer to C matrix data (M x N x L, N-major, FP16)
            shape_mnkl: Tuple of (M, N, K, L) dimensions
            max_active_clusters: Maximum active clusters for persistent scheduling
            epilogue_op: Optional epilogue operation
        """
        m, n, k, l = shape_mnkl

        # Create tensors from pointers with proper layouts
        # A is (M, K, L) K-major
        a_layout = cute.make_layout((m, k, l), stride=(k, 1, m * k))
        a_tensor = cute.make_tensor(a_ptr, a_layout)

        # B is (N, K, L) K-major
        b_layout = cute.make_layout((n, k, l), stride=(k, 1, n * k))
        b_tensor = cute.make_tensor(b_ptr, b_layout)

        # C is (M, N, L) row-major (N-stride=1, M-stride=N)
        c_layout = cute.make_layout((m, n, l), stride=(n, 1, m * n))
        c_tensor = cute.make_tensor(c_ptr, c_layout)

        # Get major modes from layouts
        self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor).mma_major_mode()
        self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor).mma_major_mode()
        self.c_layout = utils.LayoutEnum.from_tensor(c_tensor)

        # Type check
        if cutlass.const_expr(self.a_dtype != self.b_dtype):
            raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")

        # Setup kernel attributes
        self._setup_attributes()

        # Setup scale factor tensor layouts from pointers
        # Scale factors use the tile_atom_to_shape_SF layout based on A/B shapes
        sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
            a_tensor.shape, self.sf_vec_size
        )
        sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

        sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
            b_tensor.shape, self.sf_vec_size
        )
        sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)

        # Create tiled MMAs
        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype,
            self.a_major_mode,
            self.b_major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            self.cta_group,
            self.mma_inst_shape_mn,
        )

        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.a_dtype,
            self.a_major_mode,
            self.b_major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            tcgen05.CtaGroup.ONE,
            self.mma_inst_shape_mn_sfb,
        )
        atom_thr_size = cute.size(tiled_mma.thr_id.shape)

        # Setup TMA atoms for A, B, SFA, SFB, C
        a_op = sm100_utils.cluster_shape_to_tma_atom_A(
            self.cluster_shape_mn, tiled_mma.thr_id
        )
        a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
        tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
            a_op,
            a_tensor,
            a_smem_layout,
            self.mma_tiler,
            tiled_mma,
            self.cluster_layout_vmnk.shape,
        )

        b_op = sm100_utils.cluster_shape_to_tma_atom_B(
            self.cluster_shape_mn, tiled_mma.thr_id
        )
        b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
        tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
            b_op,
            b_tensor,
            b_smem_layout,
            self.mma_tiler,
            tiled_mma,
            self.cluster_layout_vmnk.shape,
        )

        sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(
            self.cluster_shape_mn, tiled_mma.thr_id
        )
        sfa_smem_layout = cute.slice_(
            self.sfa_smem_layout_staged, (None, None, None, 0)
        )
        tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
            sfa_op,
            sfa_tensor,
            sfa_smem_layout,
            self.mma_tiler,
            tiled_mma,
            self.cluster_layout_vmnk.shape,
            internal_type=cutlass.Int16,
        )

        sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
            self.cluster_shape_mn, tiled_mma.thr_id
        )
        sfb_smem_layout = cute.slice_(
            self.sfb_smem_layout_staged, (None, None, None, 0)
        )
        tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
            sfb_op,
            sfb_tensor,
            sfb_smem_layout,
            self.mma_tiler_sfb,
            tiled_mma_sfb,
            self.cluster_layout_sfb_vmnk.shape,
            internal_type=cutlass.Int16,
        )

        a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
        b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
        sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
        sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)
        self.num_tma_load_bytes = (
            a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
        ) * atom_thr_size

        # TMA store for C
        epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
        tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
            cpasync.CopyBulkTensorTileS2GOp(),
            c_tensor,
            epi_smem_layout,
            self.epi_tile,
        )

        # ---------------------------------------------------------------------
        # Deterministic persistent scheduling
        #
        # Launch a single logical cluster in (x,y) and use grid.z as a persistent
        # work queue length. All warps participate in a shared work loop and
        # synchronize once per work item.
        # ---------------------------------------------------------------------
        cluster_m, cluster_n = self.cluster_shape_mn
        v_ctas_per_tile = cute.size(tiled_mma.thr_id.shape)
        scheduler_params = SchedulerParams.from_shape(
            m=c_tensor.shape[0],
            n=c_tensor.shape[1],
            k_tile_cnt=k // self.mma_tiler[2],
            tile_m=self.mma_tiler[0] * cluster_m,
            tile_n=self.mma_tiler[1] * cluster_n,
        )

        tiles_m = scheduler_params.tiles_m
        tiles_n = scheduler_params.tiles_n
        total_work_items = scheduler_params.total_work_items(l)

        # When persistent scheduling is disabled, run one work item per cluster.
        if not self.use_persistent_scheduler:
            max_active_clusters = total_work_items

        grid_z = min(total_work_items, max_active_clusters)
        grid = (cluster_m * v_ctas_per_tile, cluster_n, grid_z)

        self.buffer_align_bytes = 1024

        # Define shared storage
        @cute.struct
        class SharedStorage:
            ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            tmem_dealloc_mbar_ptr: cutlass.Int64
            tmem_holding_buf: cutlass.Int32
            sC: cute.struct.Align[
                cute.struct.MemRange[
                    self.c_dtype, cute.cosize(self.c_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]
            sA: cute.struct.Align[
                cute.struct.MemRange[
                    self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]
            sB: cute.struct.Align[
                cute.struct.MemRange[
                    self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]
            sSFA: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]
            sSFB: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]

        self.shared_storage = SharedStorage

        self.kernel(
            tiled_mma,
            tiled_mma_sfb,
            tma_atom_a,
            tma_tensor_a,
            tma_atom_b,
            tma_tensor_b,
            tma_atom_sfa,
            tma_tensor_sfa,
            tma_atom_sfb,
            tma_tensor_sfb,
            tma_atom_c,
            tma_tensor_c,
            c_tensor,
            self.cluster_layout_vmnk,
            self.cluster_layout_sfb_vmnk,
            self.a_smem_layout_staged,
            self.b_smem_layout_staged,
            self.sfa_smem_layout_staged,
            self.sfb_smem_layout_staged,
            self.c_smem_layout_staged,
            self.epi_tile,
            tiles_m,
            tiles_n,
            l,
            total_work_items,
            epilogue_op,
        ).launch(
            grid=grid,
            block=[self.threads_per_cta, 1, 1],
            cluster=(
                self.cluster_shape_mn[0] * v_ctas_per_tile,
                self.cluster_shape_mn[1],
                1,
            ),
            min_blocks_per_mp=1,
        )

    @cute.kernel
    def kernel(
        self,
        tiled_mma: cute.TiledMma,
        tiled_mma_sfb: cute.TiledMma,
        tma_atom_a: cute.CopyAtom,
        mA_mkl: cute.Tensor,
        tma_atom_b: cute.CopyAtom,
        mB_nkl: cute.Tensor,
        tma_atom_sfa: cute.CopyAtom,
        mSFA_mkl: cute.Tensor,
        tma_atom_sfb: cute.CopyAtom,
        mSFB_nkl: cute.Tensor,
        tma_atom_c: cute.CopyAtom,
        mC_mnl: cute.Tensor,
        mC_raw_mnl: cute.Tensor,
        cluster_layout_vmnk: cute.Layout,
        cluster_layout_sfb_vmnk: cute.Layout,
        a_smem_layout_staged: cute.ComposedLayout,
        b_smem_layout_staged: cute.ComposedLayout,
        sfa_smem_layout_staged: cute.Layout,
        sfb_smem_layout_staged: cute.Layout,
        c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
        epi_tile: cute.Tile,
        tiles_m: cutlass.Int32,
        tiles_n: cutlass.Int32,
        tiles_l: cutlass.Int32,
        total_work_items: cutlass.Int32,
        epilogue_op: cutlass.Constexpr,
    ):
        """
        GPU device kernel for block-scaled GEMM.

        Kernel argument shapes (logical, as seen by the harness / custom_kernel):
        - A: `a` is (M, K_packed, L) where `K_packed = K // 2` due to FP4 packing.
             This kernel treats it as logical (M, K, L) by doubling K when needed.
        - B: `b` is (N, K_packed, L) with the same packing along K.
        - SFA (permuted): `sfa_permuted` is a CuTe/TMA-friendly view of scale factors for A.
             Reference `sfa` is (M, K//16, L); permuted view is typically shaped like
             (32, 4, rest_m, 4, rest_k, L) (exact factorization depends on tiling).
        - SFB (permuted): `sfb_permuted` is analogous for B; reference `sfb` is (N, K//16, L),
             permuted view is typically (32, 4, rest_n, 4, rest_k, L).
        - C: `c` is (M, N, L) fp16; the harness uses L as the fastest-varying dimension.
        """
        # CuTe arch intrinsics:
        # - `cute.arch.warp_idx()` gives the warp id within the CTA (0..).
        # - `cute.arch.make_warp_uniform(x)` forces a warp-uniform value (avoid divergence pitfalls).
        warp_idx = cute.arch.warp_idx()
        warp_idx = cute.arch.make_warp_uniform(warp_idx)

        # TMA descriptor prefetch:
        # - `cpasync.prefetch_descriptor(tma_atom)` hides descriptor fetch latency before first TMA issue.
        if warp_idx == self.tma_warp_id:
            cpasync.prefetch_descriptor(tma_atom_a)
            cpasync.prefetch_descriptor(tma_atom_b)
            cpasync.prefetch_descriptor(tma_atom_sfa)
            cpasync.prefetch_descriptor(tma_atom_sfb)
            cpasync.prefetch_descriptor(tma_atom_c)

        use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2

        # CTA/thread coordinates:
        # - `cute.arch.block_idx()` -> (blockIdx.x, blockIdx.y, blockIdx.z).
        # - `cute.arch.block_idx_in_cluster()` -> CTA rank within the hardware cluster.
        bidx, bidy, bidz = cute.arch.block_idx()
        v_ctas_per_tile = cute.size(tiled_mma.thr_id.shape)
        mma_tile_coord_v = bidx % v_ctas_per_tile
        is_leader_cta = mma_tile_coord_v == 0
        cta_rank_in_cluster = cute.arch.make_warp_uniform(
            cute.arch.block_idx_in_cluster()
        )
        # Cluster layout decode:
        # - `layout.get_flat_coord(rank)` maps CTA-rank -> multi-dim coords used by multicast + partitions.
        block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
            cta_rank_in_cluster
        )
        block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
            cta_rank_in_cluster
        )
        tidx, _, _ = cute.arch.thread_idx()
        # For our launch, grid.x/y enumerate CTAs within a cluster.
        cluster_n_coord = cute.arch.make_warp_uniform(bidy)

        # Shared memory allocation helper:
        # - `utils.SmemAllocator()` + `allocate(self.shared_storage)` builds a typed SMEM backing store.
        smem = utils.SmemAllocator()
        storage = smem.allocate(self.shared_storage)

        # Pipelines:
        # - `PipelineTmaUmma` coordinates TMA producer(s) and UMMA consumer(s) via mbarrier arrays.
        ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
        ab_pipeline_consumer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread, num_tma_producer
        )
        ab_pipeline = pipeline.PipelineTmaUmma.create(
            barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
            num_stages=self.num_ab_stage,
            producer_group=ab_pipeline_producer_group,
            consumer_group=ab_pipeline_consumer_group,
            tx_count=self.num_tma_load_bytes,
            cta_layout_vmnk=cluster_layout_vmnk,
            defer_sync=True,
        )

        # Accumulator pipeline:
        # - `PipelineUmmaAsync` synchronizes "accumulator stage" ownership across MMA and epilogue roles.
        acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_acc_consumer_threads = len(self.epilog_warp_id) * (
            2 if use_2cta_instrs else 1
        )
        acc_pipeline_consumer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread, num_acc_consumer_threads
        )
        acc_pipeline = pipeline.PipelineUmmaAsync.create(
            barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
            num_stages=self.num_acc_stage,
            producer_group=acc_pipeline_producer_group,
            consumer_group=acc_pipeline_consumer_group,
            cta_layout_vmnk=cluster_layout_vmnk,
            defer_sync=True,
        )

        # TMEM allocator:
        # - Provides per-CTA allocations of Blackwell tensor memory columns for accumulators + SF tiles.
        tmem = utils.TmemAllocator(
            storage.tmem_holding_buf,
            barrier_for_retrieve=self.tmem_alloc_barrier,
            allocator_warp_id=self.epilog_warp_id[0],
            is_two_cta=use_2cta_instrs,
            two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
        )

        # Cluster arrive after barrier init
        pipeline_init_arrive(cluster_shape_mn=self.cluster_shape_mn, is_relaxed=True)

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

        # TMA multicast masks:
        # - `cpasync.create_tma_multicast_mask(...)` returns the mask consumed by `cute.copy(..., mcast_mask=...)`.
        a_full_mcast_mask = None
        b_full_mcast_mask = None
        sfa_full_mcast_mask = None
        sfb_full_mcast_mask = None
        if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
            a_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )
            b_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
            )
            sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )
            sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1
            )

        # Global -> per-CTA tiles:
        # - `cute.local_tile(tensor, slice, coords)` forms a logical tile view for this CTA over GMEM tensors.
        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        gB_nkl = cute.local_tile(
            mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
        )
        gSFA_mkl = cute.local_tile(
            mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        gSFB_nkl = cute.local_tile(
            mSFB_nkl,
            cute.slice_(self.mma_tiler_sfb, (0, None, None)),
            (None, None, None),
        )
        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
        )
        # Tile sizes:
        # - `cute.size(tensor, mode=[...])` returns the extent of specific tensor mode(s) (here: K-tiles).
        k_tile_cnt = cute.size(gA_mkl, mode=[3])

        # Partition for TiledMMA
        thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
        thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
        tCgA = thr_mma.partition_A(gA_mkl)
        tCgB = thr_mma.partition_B(gB_nkl)
        tCgSFA = thr_mma.partition_A(gSFA_mkl)
        tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
        tCgC = thr_mma.partition_C(gC_mnl)

        # TMA partitions:
        # - `cpasync.tma_partition(tma_atom, cta_coord, cta_layout, smem_tile, gmem_tile)`
        #   returns (SMEM view, GMEM view) for issuing TMA copies stage-by-stage.
        a_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
        )
        tAsA, tAgA = cpasync.tma_partition(
            tma_atom_a,
            block_in_cluster_coord_vmnk[2],
            a_cta_layout,
            cute.group_modes(sA, 0, 3),
            cute.group_modes(tCgA, 0, 3),
        )

        # TMA partitions for B
        b_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
        )
        tBsB, tBgB = cpasync.tma_partition(
            tma_atom_b,
            block_in_cluster_coord_vmnk[1],
            b_cta_layout,
            cute.group_modes(sB, 0, 3),
            cute.group_modes(tCgB, 0, 3),
        )

        # TMA partitions for SFA
        sfa_cta_layout = a_cta_layout
        tAsSFA, tAgSFA = cpasync.tma_partition(
            tma_atom_sfa,
            block_in_cluster_coord_vmnk[2],
            sfa_cta_layout,
            cute.group_modes(sSFA, 0, 3),
            cute.group_modes(tCgSFA, 0, 3),
        )
        tAsSFA = cute.filter_zeros(tAsSFA)
        tAgSFA = cute.filter_zeros(tAgSFA)

        # TMA partitions for SFB
        sfb_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
        )
        tBsSFB, tBgSFB = cpasync.tma_partition(
            tma_atom_sfb,
            block_in_cluster_coord_sfb_vmnk[1],
            sfb_cta_layout,
            cute.group_modes(sSFB, 0, 3),
            cute.group_modes(tCgSFB, 0, 3),
        )
        tBsSFB = cute.filter_zeros(tBsSFB)
        tBgSFB = cute.filter_zeros(tBgSFB)

        # MMA fragment tensors
        tCrA = tiled_mma.make_fragment_A(sA)
        tCrB = tiled_mma.make_fragment_B(sB)
        acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])

        if cutlass.const_expr(self.overlapping_accum):
            num_acc_stage_overlapped = 2
            tCtAcc_fake = tiled_mma.make_fragment_C(
                cute.append(acc_shape, num_acc_stage_overlapped)
            )
            tCtAcc_fake = cute.make_tensor(
                tCtAcc_fake.iterator,
                cute.make_layout(
                    tCtAcc_fake.shape,
                    stride=(
                        tCtAcc_fake.stride[0],
                        tCtAcc_fake.stride[1],
                        tCtAcc_fake.stride[2],
                        (256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1],
                    ),
                ),
            )
        else:
            tCtAcc_fake = tiled_mma.make_fragment_C(
                cute.append(acc_shape, self.num_acc_stage)
            )

        # Wait for cluster init
        pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)

        # ---------------------------------------------------------------------
        # Deterministic persistent work loop (dev_v6 pattern)
        # ---------------------------------------------------------------------
        grid_dim_z = cute.arch.grid_dim()[2]
        grid_dim_z = cute.arch.make_warp_uniform(grid_dim_z)
        base_work_idx = cute.arch.make_warp_uniform(bidz)
        tiles_mn = tiles_m * tiles_n
        num_work = cute.ceil_div(total_work_items - base_work_idx, grid_dim_z)

        # Per-CTA pipeline states (must be defined outside dynamic control flow)
        ab_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, self.num_ab_stage
        )
        ab_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, self.num_ab_stage
        )
        acc_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, self.num_acc_stage
        )
        acc_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, self.num_acc_stage
        )

        # Allocate + retrieve TMEM pointers (once) (ported from dev_v6)
        tmem.allocate(self.num_tmem_alloc_cols)
        self.tmem_alloc_barrier.arrive_and_wait()
        tmem.wait_for_alloc()

        acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
        tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

        sfa_tmem_ptr = cute.recast_ptr(
            acc_tmem_ptr + self.num_accumulator_tmem_cols,
            dtype=self.sf_dtype,
        )
        tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
        )
        tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)

        sfb_tmem_ptr = cute.recast_ptr(
            acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
            dtype=self.sf_dtype,
        )
        tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
        )
        tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)

        (
            tiled_copy_s2t_sfa,
            tCsSFA_compact_s2t,
            tCtSFA_compact_s2t,
        ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
        (
            tiled_copy_s2t_sfb,
            tCsSFB_compact_s2t,
            tCtSFB_compact_s2t,
        ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)

        # Epilogue copy partitions (once)
        epi_tidx = tidx
        (
            tiled_copy_t2r,
            tTR_tAcc_base,
            tTR_gC_base,
            (tTR_rAcc_0, tTR_rAcc_1),
        ) = self.epilog_tmem_copy_and_partition(
            epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
        )
        tTR_rC = cute.make_rmem_tensor(tTR_rAcc_0.shape, self.c_dtype)
        tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
            tiled_copy_t2r, tTR_rC, epi_tidx, sC
        )
        (
            tma_atom_c,
            bSG_sC,
            bSG_gC_partitioned,
        ) = self.epilog_gmem_copy_and_partition(tma_atom_c, tCgC, epi_tile, sC)

        c_producer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread,
            32 * len(self.epilog_warp_id),
        )
        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=self.num_c_stage,
            producer_group=c_producer_group,
        )

        # ---------------------------------------------------------------------
        # Main work loop
        # ---------------------------------------------------------------------
        for work_iter in cutlass.range(num_work, unroll=1):
            work_idx = base_work_idx + work_iter * grid_dim_z
            # No-split schedule: work_idx -> (m,n,l)
            work_m_idx = work_idx % tiles_m
            work_n_idx = (work_idx // tiles_m) % tiles_n
            work_l_idx = work_idx // tiles_mn

            cluster_tile_coord_mnl = (work_m_idx, work_n_idx, work_l_idx)
            mma_tile_coord_mnl = (
                work_m_idx,
                work_n_idx * self.cluster_shape_mn[1] + cluster_n_coord,
                work_l_idx,
            )

            # -------------------------------------------------------------
            # TMA LOAD WARP (warp 5)
            # -------------------------------------------------------------
            if warp_idx == self.tma_warp_id:
                tAgA_slice = tAgA[
                    (None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
                ]
                tBgB_slice = tBgB[
                    (None, mma_tile_coord_mnl[1], None, cluster_tile_coord_mnl[2])
                ]
                tAgSFA_slice = tAgSFA[
                    (None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
                ]

                slice_n = cluster_tile_coord_mnl[1]
                if cutlass.const_expr(self.cluster_shape_mn[1] > 1):
                    # For clustered N, SFB needs per-CTA N tile selection (tile_n chunks).
                    slice_n = mma_tile_coord_mnl[1]
                if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
                    slice_n = slice_n // 2
                tBgSFB_slice = tBgSFB[(None, slice_n, None, cluster_tile_coord_mnl[2])]

                ab_producer_state.reset_count()
                peek_ab_empty_status = cutlass.Boolean(1)
                if ab_producer_state.count < k_tile_cnt:
                    peek_ab_empty_status = ab_pipeline.producer_try_acquire(
                        ab_producer_state
                    )

                for k_iter in cutlass.range(0, k_tile_cnt, 1, unroll=1):  # noqa: B007
                    ab_pipeline.producer_acquire(
                        ab_producer_state, peek_ab_empty_status
                    )
                    k_tile = ab_producer_state.count

                    # TMA loads:
                    # - `cute.copy(tma_atom_*, gmem_tile, smem_tile, tma_bar_ptr=..., mcast_mask=...)`
                    #   issues a TMA transaction and signals the stage mbarrier when complete.
                    cute.copy(
                        tma_atom_a,
                        tAgA_slice[(None, k_tile)],
                        tAsA[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=a_full_mcast_mask,
                    )
                    cute.copy(
                        tma_atom_b,
                        tBgB_slice[(None, k_tile)],
                        tBsB[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=b_full_mcast_mask,
                    )
                    cute.copy(
                        tma_atom_sfa,
                        tAgSFA_slice[(None, k_tile)],
                        tAsSFA[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=sfa_full_mcast_mask,
                    )
                    cute.copy(
                        tma_atom_sfb,
                        tBgSFB_slice[(None, k_tile)],
                        tBsSFB[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=sfb_full_mcast_mask,
                    )

                    ab_producer_state.advance()
                    peek_ab_empty_status = cutlass.Boolean(1)
                    if ab_producer_state.count < k_tile_cnt:
                        peek_ab_empty_status = ab_pipeline.producer_try_acquire(
                            ab_producer_state
                        )

            # -------------------------------------------------------------
            # MMA WARP (warp 4)
            # -------------------------------------------------------------
            if warp_idx == self.mma_warp_id:
                if cutlass.const_expr(self.overlapping_accum):
                    acc_stage_index = acc_producer_state.phase ^ 1
                else:
                    acc_stage_index = acc_producer_state.index

                tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]

                ab_consumer_state.reset_count()
                peek_ab_full_status = cutlass.Boolean(1)
                if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
                    peek_ab_full_status = ab_pipeline.consumer_try_wait(
                        ab_consumer_state
                    )

                if is_leader_cta:
                    acc_pipeline.producer_acquire(acc_producer_state)

                tCtSFB_mma = tCtSFB
                if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
                    offset = (
                        cutlass.Int32(2)
                        if mma_tile_coord_mnl[1] % 2 == 1
                        else cutlass.Int32(0)
                    )
                    shifted_ptr = cute.recast_ptr(
                        acc_tmem_ptr
                        + self.num_accumulator_tmem_cols
                        + self.num_sfa_tmem_cols
                        + offset,
                        dtype=self.sf_dtype,
                    )
                    tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
                elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
                    offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
                    shifted_ptr = cute.recast_ptr(
                        acc_tmem_ptr
                        + self.num_accumulator_tmem_cols
                        + self.num_sfa_tmem_cols
                        + offset,
                        dtype=self.sf_dtype,
                    )
                    tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)

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

                for k_iter in range(k_tile_cnt):  # noqa: B007
                    if is_leader_cta:
                        ab_pipeline.consumer_wait(
                            ab_consumer_state, peek_ab_full_status
                        )

                        s2t_stage_coord = (
                            None,
                            None,
                            None,
                            None,
                            ab_consumer_state.index,
                        )
                        tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
                        tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]

                        cute.copy(
                            tiled_copy_s2t_sfa,
                            tCsSFA_compact_s2t_staged,
                            tCtSFA_compact_s2t,
                        )
                        cute.copy(
                            tiled_copy_s2t_sfb,
                            tCsSFB_compact_s2t_staged,
                            tCtSFB_compact_s2t,
                        )

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

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

                            # MMA:
                            # - `cute.gemm(tiled_mma, D, A, B, C)` emits tcgen05 MMA into TMEM accumulator `D`.
                            cute.gemm(
                                tiled_mma,
                                tCtAcc,
                                tCrA[kblock_coord],
                                tCrB[kblock_coord],
                                tCtAcc,
                            )

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

                        ab_pipeline.consumer_release(ab_consumer_state)

                    ab_consumer_state.advance()
                    peek_ab_full_status = cutlass.Boolean(1)
                    if ab_consumer_state.count < k_tile_cnt:
                        if is_leader_cta:
                            peek_ab_full_status = ab_pipeline.consumer_try_wait(
                                ab_consumer_state
                            )

                if is_leader_cta:
                    acc_pipeline.producer_commit(acc_producer_state)
                acc_producer_state.advance()

            # -------------------------------------------------------------
            # EPILOGUE WARPS (warps 0-3)
            # -------------------------------------------------------------
            if warp_idx < self.mma_warp_id:
                bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]

                if cutlass.const_expr(self.overlapping_accum):
                    acc_stage_index = acc_consumer_state.phase
                else:
                    acc_stage_index = acc_consumer_state.index

                tTR_tAcc = tTR_tAcc_base[
                    (None, None, None, None, None, acc_stage_index)
                ]

                acc_pipeline.consumer_wait(acc_consumer_state)

                # Mode grouping:
                # - `cute.group_modes(t, start, end)` merges modes to produce a simpler iteration space.
                tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
                tTR_gC_tile = tTR_gC_base[
                    (
                        None,
                        None,
                        None,
                        None,
                        None,
                        mma_tile_coord_mnl[0],
                        mma_tile_coord_mnl[1],
                        mma_tile_coord_mnl[2],
                    )
                ]
                tTR_gC = cute.group_modes(tTR_gC_tile, 3, cute.rank(tTR_gC_tile))
                bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))

                subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
                num_prev_subtiles = 0

                if warp_idx == self.epilog_warp_id[0] and subtile_cnt > 0:
                    c_pipeline.producer_acquire()

                if subtile_cnt > 0:
                    tTR_tAcc_mn_first = tTR_tAcc[(None, None, None, 0)]
                    cute.copy(tiled_copy_t2r, tTR_tAcc_mn_first, tTR_rAcc_0)

                for subtile_idx in range(subtile_cnt):
                    curr_buf = tTR_rAcc_0 if subtile_idx % 2 == 0 else tTR_rAcc_1
                    next_buf = tTR_rAcc_1 if subtile_idx % 2 == 0 else tTR_rAcc_0

                    if subtile_idx + 1 < subtile_cnt:
                        tTR_tAcc_mn_next = tTR_tAcc[(None, None, None, subtile_idx + 1)]
                        cute.copy(tiled_copy_t2r, tTR_tAcc_mn_next, next_buf)

                    acc_vec = tiled_copy_r2s.retile(curr_buf).load()
                    tRS_rC.store(epilogue_op(acc_vec).to(self.c_dtype))

                    c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
                    cute.copy(
                        tiled_copy_r2s,
                        tRS_rC,
                        tRS_sC[(None, None, None, c_buffer)],
                    )

                    # Memory ordering for TMA store:
                    # - `cute.arch.fence_proxy(async_shared)` makes prior SMEM writes visible to async consumers (TMA).
                    cute.arch.fence_proxy(
                        cute.arch.ProxyKind.async_shared,
                        space=cute.arch.SharedSpace.shared_cta,
                    )
                    self.epilog_sync_barrier.arrive_and_wait()

                    if warp_idx == self.epilog_warp_id[0]:
                        cute.copy(
                            tma_atom_c,
                            bSG_sC[(None, c_buffer)],
                            bSG_gC[(None, subtile_idx)],
                        )
                        c_pipeline.producer_commit()
                        if (subtile_idx + 1) < subtile_cnt:
                            c_pipeline.producer_acquire()
                    self.epilog_sync_barrier.arrive_and_wait()

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

            # Keep warps/roles in lockstep across work items.
            self.work_loop_barrier.arrive_and_wait()

        # ---------------------------------------------------------------------
        # Tail / teardown
        # ---------------------------------------------------------------------
        if warp_idx == self.tma_warp_id:
            ab_pipeline.producer_tail(ab_producer_state)

        if warp_idx == self.mma_warp_id:
            acc_pipeline.producer_tail(acc_producer_state)

        if warp_idx < self.mma_warp_id:
            # Wait for C TMA store complete
            c_pipeline.producer_tail()

            # Free TMEM
            tmem.relinquish_alloc_permit()
            self.epilog_sync_barrier.arrive_and_wait()
            tmem.free(acc_tmem_ptr)

    def mainloop_s2t_copy_and_partition(self, sSF, tCtSF):
        """Create S2T copy for scale factors.

        Uses the tcgen05 S2T (SMEM to TMEM) copy operations for Blackwell.
        Following the reference pattern from grouped_blockscaled_gemm.py.
        """
        # Filter zeros from both source and destination tensors
        # (MMA, MMA_MN, MMA_K, STAGE)
        tCsSF_compact = cute.filter_zeros(sSF)
        # (MMA, MMA_MN, MMA_K)
        tCtSF_compact = cute.filter_zeros(tCtSF)

        # Make S2T CopyAtom using Cp4x32x128bOp (32x128b with warpx4 broadcast)
        copy_atom_s2t = cute.make_copy_atom(
            tcgen05.Cp4x32x128bOp(self.cta_group),
            self.sf_dtype,
        )
        # Create the tiled copy using the tmem tensor
        tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
        thr_copy_s2t = tiled_copy_s2t.get_slice(0)

        # Partition source (SMEM) and get descriptors
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
        tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
        tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t, tCsSF_compact_s2t_
        )
        # Partition destination (TMEM)
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
        tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)

        return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t

    def epilog_tmem_copy_and_partition(
        self, tidx, tAcc, gC_mnl, epi_tile, use_2cta_instrs
    ):
        """Partition for TMEM to register copy in epilogue.

        Uses sm100_utils.get_tmem_load_op and tcgen05.make_tmem_copy to create
        the tiled copy for TMEM to register.
        """
        # Make tiledCopy for tensor memory load
        copy_atom_t2r = sm100_utils.get_tmem_load_op(
            self.cta_tile_shape_mnk,
            self.c_layout,
            self.c_dtype,
            self.acc_dtype,
            epi_tile,
            use_2cta_instrs,
        )
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE)
        tAcc_epi = cute.flat_divide(
            tAcc[((None, None), 0, 0, None)],
            epi_tile,
        )
        # (EPI_TILE_M, EPI_TILE_N)
        tiled_copy_t2r = tcgen05.make_tmem_copy(
            copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
        )

        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
        # (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE)
        tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)

        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_mnl_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )
        # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        # (T2R, T2R_M, T2R_N)
        # Double-buffer for TMEM prefetch optimization
        tTR_rAcc_0 = cute.make_rmem_tensor(
            tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
        )
        tTR_rAcc_1 = cute.make_rmem_tensor(
            tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
        )

        return tiled_copy_t2r, tTR_tAcc, tTR_gC, (tTR_rAcc_0, tTR_rAcc_1)

    def epilog_smem_copy_and_partition(self, tiled_copy_t2r, tTR_rC, tidx, sC):
        """Partition for register to SMEM copy in epilogue.

        Uses sm100_utils.get_smem_store_op and cute.make_tiled_copy_D.
        """
        copy_atom_r2s = sm100_utils.get_smem_store_op(
            self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
        )
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
        # (R2S, R2S_M, R2S_N, PIPE_D)
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        # (R2S, R2S_M, R2S_N)
        tRS_rC = tiled_copy_r2s.retile(tTR_rC)

        return tiled_copy_r2s, tRS_rC, tRS_sC

    def epilog_gmem_copy_and_partition(self, tma_atom_c, gC_mnl, epi_tile, sC):
        """Partition for SMEM to GMEM TMA store in epilogue.

        Uses cpasync.tma_partition to partition both source (SMEM) and destination (GMEM).
        """
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )

        sC_for_tma_partition = cute.group_modes(sC, 0, 2)
        gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
        # ((ATOM_V, REST_V), EPI_M, EPI_N)
        # ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
        bSG_sC, bSG_gC = cpasync.tma_partition(
            tma_atom_c,
            0,
            cute.make_layout(1),
            sC_for_tma_partition,
            gC_for_tma_partition,
        )

        return tma_atom_c, bSG_sC, bSG_gC


# Global compiled kernel cache
# - key: (m, n, k, l, mma_tiler_mn, cluster_shape_mn, num_ab_stage, num_c_stage)
_compiled_kernel_cache = {}

# Data type constants
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16


# Feature flags
USE_PERSISTENT_SCHEDULER = True  # Set to False to disable persistent scheduling

# Optional compilation debug flags (off by default).
# Enable with `CUTE_LINEINFO=1` to generate line info in the compiled cubin/ptx.
_ENABLE_CUTE_LINEINFO = os.getenv("CUTE_LINEINFO", "0") == "1"


# Optimal configuration lookup: maps (m, n, k, l) -> (mma_tiler_mn, cluster_shape_mn)
# Enable (1, 2) clustering to use TMA multicast where applicable.
@dataclass(frozen=True)
class KernelConfig:
    mma_tiler_mn: Tuple[int, int]
    cluster_shape_mn: Tuple[int, int]
    num_ab_stage: int = 0  # 0 => auto (use _compute_stages())
    num_c_stage: int = 0  # 0 => auto (use _compute_stages())


_OPTIMAL_CONFIG = {
    (128, 7168, 16384, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2
    ),
    (128, 4096, 7168, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2
    ),
    (128, 7168, 2048, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2
    ),
}

# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
    (128, 64),
    (1, 1),
    num_ab_stage=0,
    num_c_stage=0,
)


def _env_override_int(name: str) -> Optional[int]:
    val = os.getenv(name)
    if val is None or val == "":
        return None
    return int(val)


def _cache_key(
    m: int,
    n: int,
    k: int,
    l: int,
    mma_tiler_mn: Tuple[int, int],
    cluster_shape_mn: Tuple[int, int],
    num_ab_stage: int,
    num_c_stage: int,
) -> Tuple[int, int, int, int, Tuple[int, int], Tuple[int, int], int, int]:
    return (
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        num_ab_stage,
        num_c_stage,
    )


def _ensure_compiled(
    m: int,
    n: int,
    k: int,
    l: int,
    mma_tiler_mn: Tuple[int, int],
    cluster_shape_mn: Tuple[int, int],
    num_ab_stage: int,
    num_c_stage: int,
) -> None:
    """Compile a single (shape, tiling) variant if it's missing."""
    global _compiled_kernel_cache
    key = _cache_key(
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        num_ab_stage,
        num_c_stage,
    )
    if key in _compiled_kernel_cache:
        return
    num_ab_stage_override = None if num_ab_stage <= 0 else num_ab_stage
    num_c_stage_override = None if num_c_stage <= 0 else num_c_stage

    kernel_instance = Sm100BlockScaledGemmKernel(
        sf_vec_size=16,
        mma_tiler_mn=mma_tiler_mn,
        cluster_shape_mn=cluster_shape_mn,
        num_ab_stage=num_ab_stage_override,
        num_c_stage=num_c_stage_override,
        ab_dtype=ab_dtype,
        sf_dtype=sf_dtype,
        c_dtype=c_dtype,
        use_persistent_scheduler=USE_PERSISTENT_SCHEDULER,
    )

    # Placeholder pointers for compilation
    a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    b_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)
    sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)

    compile_fn = (
        cute.compile[cute.GenerateLineInfo(True)]
        if _ENABLE_CUTE_LINEINFO
        else cute.compile
    )

    _compiled_kernel_cache[key] = compile_fn(
        kernel_instance,
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
        (m, n, k, l),
    )


def compile_kernel():
    """
    Pre-compile ALL required kernel configurations.
    This is called by competition code before benchmarking to warm up.
    """
    global _compiled_kernel_cache

    # Some tests intentionally set `_compiled_kernel_cache = None` to force a cold compile.
    # Re-initialize here so both `compile_kernel()` and `custom_kernel()` stay robust.
    if _compiled_kernel_cache is None:
        _compiled_kernel_cache = {}

    # Precompile per-problem variants for the known benchmark shapes.
    for shape_key, cfg in _OPTIMAL_CONFIG.items():
        m, n, k, l = shape_key
        num_ab_stage = max(0, int(cfg.num_ab_stage))
        num_c_stage = max(0, int(cfg.num_c_stage))
        _ensure_compiled(
            m,
            n,
            k,
            l,
            cfg.mma_tiler_mn,
            cfg.cluster_shape_mn,
            num_ab_stage,
            num_c_stage,
        )

    # Conservative fallback so cold shapes still work.
    _ensure_compiled(
        128,
        256,
        256,
        1,
        _DEFAULT_CONFIG.mma_tiler_mn,
        _DEFAULT_CONFIG.cluster_shape_mn,
        _DEFAULT_CONFIG.num_ab_stage,
        _DEFAULT_CONFIG.num_c_stage,
    )


# Cache uses default arg trick: mutable default is evaluated once at definition,
# stored in function object, accessed via LOAD_FAST (faster than LOAD_GLOBAL)
_gmem = cute.AddressSpace.gmem


def custom_kernel(data: input_t, _c=[None] * 7) -> output_t:
    """
    Execute kernel with pre-compiled configuration from cache.

    Cache structure (_c):
    - _c[0]: compiled kernel for current key
    - _c[1-5]: cached pointer objects (a, b, sfa, sfb, c)
    - _c[6]: last cache_key
    """
    a, b, _, _, sfa_p, sfb_p, c = data

    global _compiled_kernel_cache
    if _compiled_kernel_cache is None:
        _compiled_kernel_cache = {}

    # Get input dimensions
    m, n, l = c.shape
    k = a.shape[1] << 1  # K is doubled due to FP4 packing

    # Get optimal config for this input shape via lookup
    shape_key = (m, n, k, l)
    cfg = _OPTIMAL_CONFIG.get(shape_key, _DEFAULT_CONFIG)
    mma_tiler_mn = cfg.mma_tiler_mn
    cluster_shape_mn = cfg.cluster_shape_mn
    num_ab_stage = _env_override_int("AB_STAGES")
    num_ab_stage = int(cfg.num_ab_stage) if num_ab_stage is None else num_ab_stage
    if num_ab_stage < 0:
        num_ab_stage = 0
    num_c_stage = _env_override_int("C_STAGES")
    num_c_stage = int(cfg.num_c_stage) if num_c_stage is None else num_c_stage
    if num_c_stage < 0:
        num_c_stage = 0
    cache_key = _cache_key(
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        num_ab_stage,
        num_c_stage,
    )

    # Initialize pointer cache on first call (independent of kernel config)
    if _c[1] is None:
        _c[1] = make_ptr(ab_dtype, 16, _gmem, assumed_align=16)
        _c[2] = make_ptr(ab_dtype, 16, _gmem, assumed_align=16)
        _c[3] = make_ptr(sf_dtype, 32, _gmem, assumed_align=32)
        _c[4] = make_ptr(sf_dtype, 32, _gmem, assumed_align=32)
        _c[5] = make_ptr(c_dtype, 16, _gmem, assumed_align=16)

    # Swap in compiled kernel when cache key changes
    if _c[6] != cache_key or _c[0] is None:
        if cache_key not in _compiled_kernel_cache:
            _ensure_compiled(
                m,
                n,
                k,
                l,
                mma_tiler_mn,
                cluster_shape_mn,
                num_ab_stage,
                num_c_stage,
            )
        _c[0] = _compiled_kernel_cache[cache_key]
        _c[6] = cache_key

    # Update pointer addresses
    _c[1]._pointer = a.data_ptr()
    _c[1]._c_pointer = None
    _c[2]._pointer = b.data_ptr()
    _c[2]._c_pointer = None
    _c[3]._pointer = sfa_p.data_ptr()
    _c[3]._c_pointer = None
    _c[4]._pointer = sfb_p.data_ptr()
    _c[4]._c_pointer = None
    _c[5]._pointer = c.data_ptr()
    _c[5]._c_pointer = None

    # Launch pre-compiled kernel
    _c[0](
        _c[1],
        _c[2],
        _c[3],
        _c[4],
        _c[5],
        (m, n, k, l),
    )
    return c
scrolls · 1725 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 154166.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON