Skip to content
KernelIndex
Search⌘K

submission 154166

rcmalli · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-154166?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
13.3µs
#119 of 369
2025-12-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5c317850a9296ad148ccaa25502418d652ab59476a2cc542ac7ab9538b239d58
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
split-kworkspace_ptr: cute.Pointer, # Workspace for Split-K reduction
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.py2798 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 torch

import cutlass
import cutlass.cute as cute
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import llvm
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


# ============================================================================
# dev_v8 developer guidance (comments-only)
# ----------------------------------------------------------------------------
# This file is intended to be modified by the human author. The agent may add
# comments to guide the strm-K / Split-K integration, but should not change
# kernel behavior unless explicitly requested.
#
# Start here:
# - Plan: `dev_v8/PLAN_strmK_SPLITK.md`
# - OpenSpec notes: `openspec/changes/plan-strmk-splitk-dev-v8/`
# - CuTeDSL API list: `dev_v8/CUTEDSL_API_REFERENCE.md`
#
# Primary goal for strm-K/Split-K:
# - Increase parallelism for small-M / large-K shapes (e.g. m=128,n=7168,k=16384)
#   where NCU shows a small grid and low waves/SM. strm-K introduces additional
#   work items per (m,n,l) tile by splitting the K tiles.
# ============================================================================


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

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

    def total_work_items(self, l: int) -> int:
        # DSL-friendly: `split_factor` is the sole control for strm-K/Split-K work expansion.
        return self.tiles_mn * l * self.split_factor

    @classmethod
    def from_shape(
        cls,
        m: int,
        n: int,
        k_tile_cnt: int,
        tile_m: int,
        tile_n: int,
        split_factor: 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, split_factor)


@dsl_user_op
def red_global_add_noftz_f16(ptr: cute.Pointer, val, *, loc=None, ip=None) -> None:
    """Atomic reduction add into global memory for FP16 using `red.global.add.noftz.f16`."""
    val_bits_i16 = llvm.bitcast(T.i16(), val.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
    llvm.inline_asm(
        None,
        [
            ptr.toint(loc=loc, ip=ip).ir_value(loc=loc, ip=ip),
            val_bits_i16,
        ],
        "red.global.add.noftz.f16 [$0], $1;",
        "l,h",
        has_side_effects=True,
        asm_dialect=0,
        loc=loc,
        ip=ip,
    )


def _leaf0(x):
    """Extract the first leaf element from a potential XTuple/tuple."""
    while cute.depth(x) > 0:
        x = cute.get(x, mode=[0])
    return x


def _leaf_last(x):
    """Extract the last leaf element from a potential XTuple/tuple."""
    while cute.depth(x) > 0:
        x = cute.get(x, mode=[cute.rank(x) - 1])
    return x


def _leaf_ptr(x):
    """Extract a `cute.Pointer` leaf from a potential XTuple/tuple."""
    a = _leaf0(x)
    if hasattr(a, "toint"):
        return a
    b = _leaf_last(x)
    if hasattr(b, "toint"):
        return b
    return a


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,
        workspace_ptr: cute.Pointer,  # Workspace for Split-K reduction
        shape_mnkl: Tuple[int, int, int, int],
        split_factor: cutlass.Int32 = 1,
        max_active_clusters: cutlass.Constexpr = 132,
        epilogue_op: cutlass.Constexpr = lambda x: x,
        use_atomic_epilogue: cutlass.Constexpr = False,
    ):
        """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)
            workspace_ptr: Pointer to workspace for Split-K partial sums (M*N*L*(split_factor-1), FP16)
            shape_mnkl: Tuple of (M, N, K, L) dimensions
            max_active_clusters: Maximum active clusters for persistent scheduling
            epilogue_op: Optional epilogue operation
            split_factor: Runtime K-split factor (1 disables strm-K/Split-K)
        """
        m, n, k, l = shape_mnkl

        # strm-K / SPLIT-K HOOK (host-side planning):
        # - Define strm-K scheduler params (tiles_m, tiles_n, tiles_k, split_factor, etc.)
        #   and decide whether to run data-parallel vs split-K/strm-K for this shape.
        # - Keep this computation in Python/host path when possible, then pass the params
        #   into `kernel(...)` as extra arguments.
        # - Reference: `dev_v8/PLAN_strmK_SPLITK.md` (Sections 1, 4) and
        #   `.claude/skills/cute-dsl/cute-dsl-strm-k-scheduling.md`.

        # Data types are set during __init__ (self.a_dtype, self.b_dtype, etc.)
        # as JIT-compiled code cannot extract element_type from _Pointer objects

        # 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)
        # PyTorch tensor after permute(1,2,0) has strides (N, 1, M*N)
        c_layout = cute.make_layout((m, n, l), stride=(n, 1, m * n))
        c_tensor = cute.make_tensor(c_ptr, c_layout)

        # Workspace for Split-K reduction: flat buffer of size M*N*L*(split_factor-1)
        # Each split partition (except the last) writes its partial sum here
        # Layout is flat (1D) - we index by: g_idx + k_part * (M*N*L)
        workspace_size = m * n * l * (split_factor - 1) if split_factor > 1 else 1
        workspace_layout = cute.make_layout((workspace_size,), stride=(1,))
        workspace_tensor = cute.make_tensor(workspace_ptr, workspace_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 (ported from dev_v6)
        #
        # 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)
        # strm-K / Split-K control:
        # `split_factor` is a runtime argument so we can tune without recompiling.
        # NOTE: Split-K requires a reduction strategy in the epilogue (not implemented yet),
        # so callers should keep `split_factor=1` until that work lands.
        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,
            split_factor=split_factor,
        )

        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

        # Launch kernel (two isolated entrypoints)
        if cutlass.const_expr(use_atomic_epilogue):
            self.kernel_split(
                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,
                workspace_tensor,  # Workspace for Split-K reduction
                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,
                split_factor=split_factor,
            ).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,
            )
        else:
            self.kernel_nosplit(
                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.jit
    def call_tma(
        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,
    ):
        # Dedicated TMA epilogue launcher (no Split-K).
        # Note: workspace_ptr is not used when split_factor=1, but we need to pass
        # a valid pointer to satisfy __call__ signature. Use c_ptr as dummy.
        return self(
            a_ptr,
            b_ptr,
            sfa_ptr,
            sfb_ptr,
            c_ptr,
            c_ptr,  # workspace_ptr (unused, split_factor=1)
            shape_mnkl,
            cutlass.Int32(1),
            max_active_clusters=max_active_clusters,
            epilogue_op=epilogue_op,
        )

    @cute.jit
    def call_atomic(
        self,
        a_ptr: cute.Pointer,
        b_ptr: cute.Pointer,
        sfa_ptr: cute.Pointer,
        sfb_ptr: cute.Pointer,
        c_ptr: cute.Pointer,
        workspace_ptr: cute.Pointer,  # Workspace for Split-K reduction
        shape_mnkl: Tuple[int, int, int, int],
        split_factor: cutlass.Constexpr,
        max_active_clusters: cutlass.Constexpr = 132,
        epilogue_op: cutlass.Constexpr = lambda x: x,
    ):
        # Dedicated workspace reduction launcher (Split-K / strm-K partial sums).
        return self(
            a_ptr,
            b_ptr,
            sfa_ptr,
            sfb_ptr,
            c_ptr,
            workspace_ptr,
            shape_mnkl,
            split_factor,
            max_active_clusters=max_active_clusters,
            epilogue_op=epilogue_op,
        )

    @staticmethod
    @cute.jit
    def decode_work_idx(
        work_idx,
        tiles_m,
        tiles_n,
        tiles_mn,
        k_tile_cnt,
        split_factor: cutlass.Constexpr,
    ):
        work_m_idx = work_idx % tiles_m
        work_n_idx = (work_idx // tiles_m) % tiles_n
        work_l_idx = work_idx // tiles_mn

        # strm-K split along K tiles (even partition; empty ranges allowed if split_factor > k_tile_cnt)
        k_part = (work_idx // tiles_mn) % split_factor
        k_tile_start = (k_part * k_tile_cnt) // split_factor
        k_tile_end = ((k_part + 1) * k_tile_cnt) // split_factor
        # Return k_part for workspace reduction (split index: 0..split_factor-1)
        return work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, k_part

    @cute.kernel
    def kernel_nosplit(
        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)

    @cute.kernel
    def kernel_split(
        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,
        mWorkspace_mnl: cute.Tensor,  # Workspace tensor: (M, N, L, split_factor-1)
        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,
        split_factor: cutlass.Constexpr = 1,
    ):
        """
        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.
        """
        use_atomic_epilogue = True

        # 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
        )
        if cutlass.const_expr(not use_atomic_epilogue):
            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
            # strm-K / SPLIT-K HOOK (device-side decode):
            # Replace the simple (m,n,l) decode below with a strm-K decode that returns:
            #   (work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, is_splitk)
            #
            # Minimal first step (no behavior change):
            # - Implement decode functions but keep `k_tile_start=0`, `k_tile_end=k_tile_cnt`,
            #   `is_splitk=False` so correctness stays identical.
            #
            # Full strm-K step:
            # - Use `k_tile_start/k_tile_end` to restrict the K-tile loop for both TMA and MMA.
            # - Ensure K coverage invariant holds (see plan + `.claude/.../cute-dsl-kernel-testing.md`).

            work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, k_part = (
                self.decode_work_idx(
                    work_idx,
                    tiles_m,
                    tiles_n,
                    tiles_mn,
                    k_tile_cnt,
                    split_factor,
                )
            )
            k_tile_cnt_local = k_tile_end - k_tile_start
            # Workspace reduction: last split writes to output, others to workspace
            is_last_split = k_part == split_factor - 1

            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_local:
                    peek_ab_empty_status = ab_pipeline.producer_try_acquire(
                        ab_producer_state
                    )

                # strm-K / SPLIT-K HOOK:
                # For split-K, change the loop to iterate only over the assigned K tile range:
                #   for _ in range(k_tile_cnt_local):
                #     k_tile = k_tile_start + ab_producer_state.count
                # and index gmem tiles using `k_tile` (not raw `count`).
                for k_iter in cutlass.range(0, k_tile_cnt_local, 1, unroll=1):  # noqa: B007
                    ab_pipeline.producer_acquire(
                        ab_producer_state, peek_ab_empty_status
                    )
                    k_tile = k_tile_start + 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_local:
                        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_local 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)

                # strm-K / SPLIT-K HOOK:
                # For split-K, iterate only the assigned K-tile subset.
                # IMPORTANT: You still want `Field.ACCUMULATE=False` for the first kblock of the
                # first K tile in THIS CTA's range, then `True` for subsequent kblocks/K-tiles.
                for k_iter in range(k_tile_cnt_local):  # 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_local:
                        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:
                # strm-K / SPLIT-K HOOK:
                # For split-K tiles, this epilogue cannot directly TMA-store to `C` because multiple
                # CTAs contribute partial sums to the same output tile.
                #
                # Typical workspace-free strategy: convert fragments to FP16/FP32 and use global atomics
                # to accumulate into `C`. This requires guaranteeing `C` is initialized to zero for
                # split-K tiles (or implementing a safe init protocol).
                #
                # Reference:
                # - `dev_v8/PLAN_strmK_SPLITK.md` (Output semantics)
                # - `.claude/skills/cute-dsl/cute-dsl-nvfp4-gemv-optimization.md` (nvvm.atomicrmw FADD pattern)
                # - `dev_v8/CUTEDSL_API_REFERENCE.md` (atomics gap section)
                if cutlass.const_expr(not use_atomic_epilogue):
                    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))
                # Slice to the current output tile so coordinate-based addressing works for Split-K atomics.
                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))
                if cutlass.const_expr(not use_atomic_epilogue):
                    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 cutlass.const_expr(not use_atomic_epilogue):
                    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)

                    if cutlass.const_expr(use_atomic_epilogue):
                        # Split-K epilogue with in-kernel workspace reduction (no atomics).
                        # - Splits 0 to (N-2): write partials to workspace[k_part]
                        # - Split (N-1): read all workspace partials, sum with own partial, write to C
                        # This eliminates the need for a separate reduction kernel.
                        tTR_gC_mn = tTR_gC[(None, None, None, subtile_idx)]
                        elem_cnt = cute.size(curr_buf.shape)
                        output_size = cute.size(mC_raw_mnl.layout)

                        # Element-wise processing (no atomics) - hardware will coalesce
                        for elem_i in cutlass.range(elem_cnt, unroll_full=True):
                            crd = cute.idx2crd(elem_i, curr_buf.shape)
                            val_acc = curr_buf[crd]
                            val_c = epilogue_op(val_acc).to(self.c_dtype)
                            g_crd = tTR_gC_mn[crd]
                            g_idx = cute.crd2idx(g_crd, mC_raw_mnl.layout)

                            if is_last_split:
                                # Last split: read all workspace partials, sum, write final result
                                # Accumulate partials from all previous splits
                                total = val_c
                                for p in range(split_factor - 1):
                                    workspace_val = mWorkspace_mnl[
                                        g_idx + p * output_size
                                    ]
                                    total = total + workspace_val
                                mC_raw_mnl[g_idx] = total
                            else:
                                # Other splits: write to workspace slice
                                # Workspace offset: k_part * (M*N*L)
                                mWorkspace_mnl[g_idx + k_part * output_size] = val_c
                    else:
                        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:
            if cutlass.const_expr(not use_atomic_epilogue):
                # 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, split_factor)
_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]
    split_factor: int = 1
    num_ab_stage: int = 0  # 0 => auto (use _compute_stages())
    num_c_stage: int = 0  # 0 => auto (use _compute_stages())


_OPTIMAL_CONFIG = {
    # You manage all tuning here per problem type:
    # KernelConfig(mma_tiler_mn=(tile_m, tile_n), cluster_shape_mn=(cluster_m, cluster_n),
    #             split_factor=..., num_ab_stage=...)
    (128, 7168, 16384, 1): KernelConfig(
        (128, 128), (1, 2), split_factor=1, num_ab_stage=0
    ),
    (128, 4096, 7168, 1): KernelConfig(
        (128, 64), (1, 2), split_factor=1, num_ab_stage=0
    ),
    (128, 7168, 2048, 1): KernelConfig(
        (128, 64), (1, 2), split_factor=1, num_ab_stage=0
    ),
}

# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
    (128, 64),
    (1, 1),
    split_factor=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],
    split_factor: int,
    num_ab_stage: int,
    num_c_stage: int,
) -> Tuple[int, int, int, int, Tuple[int, int], Tuple[int, int], int, int, int]:
    return (
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        split_factor,
        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],
    split_factor: int,
    num_ab_stage: int,
    num_c_stage: int,
) -> None:
    """Compile a single (shape, tiling, split_factor) variant if it's missing."""
    global _compiled_kernel_cache
    key = _cache_key(
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        split_factor,
        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)
    workspace_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)

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

    if split_factor > 1:
        _compiled_kernel_cache[key] = compile_fn(
            kernel_instance.call_atomic,
            a_ptr,
            b_ptr,
            sfa_ptr,
            sfb_ptr,
            c_ptr,
            workspace_ptr,
            (m, n, k, l),
            split_factor,
        )
    else:
        _compiled_kernel_cache[key] = compile_fn(
            kernel_instance.call_tma,
            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
        split_factor = max(1, int(cfg.split_factor))
        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,
            split_factor,
            num_ab_stage,
            num_c_stage,
        )

    # Conservative fallback (split_factor=1) so cold shapes still work.
    _ensure_compiled(
        128,
        256,
        256,
        1,
        _DEFAULT_CONFIG.mma_tiler_mn,
        _DEFAULT_CONFIG.cluster_shape_mn,
        _DEFAULT_CONFIG.split_factor,
        _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] * 9) -> 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
    - _c[7]: workspace pointer
    - _c[8]: workspace tensor (reused across calls)
    """
    a, b, _, _, sfa_p, sfb_p, c = data

    global _compiled_kernel_cache
    if _compiled_kernel_cache is None:
        _compiled_kernel_cache = {}

    # -------------------------------------------------------------------------
    # Input / output tensor shapes from the task harness:
    #
    # Inputs:
    # - `a`: (M, K_packed, L) nvfp4(e2m1), K-major. NOTE: `K_packed = K//2` (FP4 packed).
    # - `b`: (N, K_packed, L) nvfp4(e2m1), K-major. NOTE: `K_packed = K//2` (FP4 packed).
    # - `sfa`: (M, K//16, L) fp8(e4m3fnuz) reference layout (NOT used by this kernel).
    # - `sfb`: (N, K//16, L) fp8(e4m3fnuz) reference layout (NOT used by this kernel).
    # - `sfa_p` (`sfa_permuted`): CuTe kernel scale-factor layout for A, typically
    #   shaped like (32, 4, rest_m, 4, rest_k, L) (exact factorization depends on tiling).
    # - `sfb_p` (`sfb_permuted`): CuTe kernel scale-factor layout for B, typically
    #   shaped like (32, 4, rest_n, 4, rest_k, L).
    #
    # Output:
    # - `c`: (M, N, L) fp16 output tile accumulator / destination.
    # -------------------------------------------------------------------------

    # 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
    # strm-K / Split-K tuning parameter.
    # NOTE: `split_factor` is compile-time-specialized (part of the cache key) for max performance.
    # Changing it at runtime selects/compiles a different cached kernel variant.
    split_factor = _env_override_int("SPLIT_FACTOR")
    split_factor = int(cfg.split_factor) if split_factor is None else split_factor
    if split_factor < 1:
        split_factor = 1
    # Note: With workspace reduction, C does NOT need to be zero-initialized.
    # The last split writes directly to C (non-atomic), and workspace partials are summed.
    use_atomic_epilogue = split_factor > 1
    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,
        split_factor,
        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)
        _c[7] = make_ptr(c_dtype, 16, _gmem, assumed_align=16)  # workspace pointer

    # 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,
                split_factor,
                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
    if use_atomic_epilogue:
        # Workspace reduction path: allocate workspace, run GEMM kernel
        # The kernel performs in-kernel reduction: last split reads workspace partials,
        # sums them with its own partial, and writes the final result to C.
        # No separate reduction kernel needed!
        workspace_size = m * n * l * (split_factor - 1)

        # Allocate or reuse workspace tensor
        if _c[8] is None or _c[8].numel() < workspace_size:
            _c[8] = torch.zeros(workspace_size, dtype=torch.float16, device=c.device)
        workspace = _c[8][:workspace_size]

        # Update workspace pointer
        _c[7]._pointer = workspace.data_ptr()
        _c[7]._c_pointer = None

        # Run GEMM kernel - in-kernel reduction handles everything
        _c[0](_c[1], _c[2], _c[3], _c[4], _c[5], _c[7], (m, n, k, l))
    else:
        _c[0](
            _c[1],
            _c[2],
            _c[3],
            _c[4],
            _c[5],
            (m, n, k, l),
        )
    return c
scrolls · 2798 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 126536.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON