Skip to content
KernelIndex
Search⌘K

submission 179783

rcmalli · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9a3a9081535360164ebece48eb3dfcc2dd9def9389166f9852b9324393234343
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- Independent TMA/MMA/Epilogue work loops (no global sync barrier)
mbarrierself.epilog_sync_barrier = pipeline.NamedBarrier(
persistent-kernelThis implements a persistent warp-specialized kernel for B200 (SM100) targeting
shared-memorydef make_smem_layout_sfa_custom(
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.py1793 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.

V10: Clean rewrite with independent work loops per warp role
- Independent TMA/MMA/Epilogue work loops (no global sync barrier)
- StaticPersistentTileScheduler per warp role
- TMEM allocation only for MMA and Epilogue warps
- Proper barrier scoping
"""

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

import os

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

from task import input_t, output_t


# =============================================================================
# Custom SMEM Layout Functions for Scale Factors
# These are modified versions of blockscaled_utils functions that accept
# mma_inst_tile_k as a parameter instead of hardcoding it to 4.
# =============================================================================


@dataclass(frozen=True)
class BlockScaledBasicChunk:
    """
    The basic scale factor atom layout decided by tcgen05 BlockScaled MMA Ops.
    Copied from cutlass.utils.blockscaled_layout to avoid import issues.
    """

    sf_vec_size: int

    @property
    def layout(self) -> cute.Layout:
        # K-major layout: (AtomMN, AtomK)
        atom_shape = ((32, 4), (self.sf_vec_size, 4))
        atom_stride = ((16, 4), (0, 1))
        return cute.make_layout(atom_shape, stride=atom_stride)


def make_smem_layout_sfa_custom(
    tiled_mma: cute.TiledMma,
    mma_tiler_mnk: Tuple[int, int, int],
    sf_vec_size: int,
    num_stages: int,
    mma_inst_tile_k: int = 4,  # NEW: Configurable parameter
) -> cute.Layout:
    """
    Custom version of blockscaled_utils.make_smem_layout_sfa that accepts
    mma_inst_tile_k as a parameter instead of hardcoding it to 4.
    """
    # (CTA_Tile_Shape_M, MMA_Tile_Shape_K)
    sfa_tile_shape = (
        mma_tiler_mnk[0] // cute.size(tiled_mma.thr_id.shape),
        mma_tiler_mnk[2],
    )

    # ((Atom_M, Rest_M),(Atom_K, Rest_K))
    smem_layout = cute.tile_to_shape(
        BlockScaledBasicChunk(sf_vec_size).layout,
        sfa_tile_shape,
        (2, 1),
    )

    # Use the passed mma_inst_tile_k instead of hardcoded 4
    # (CTA_Tile_Shape_M, MMA_Inst_Shape_K)
    sfa_tile_shape = cute.shape_div(sfa_tile_shape, (1, mma_inst_tile_k))
    # ((Atom_Inst_M, Atom_Inst_K), MMA_M, MMA_K))
    smem_layout = cute.tiled_divide(smem_layout, sfa_tile_shape)

    atom_m = 128
    tiler_inst = ((atom_m, sf_vec_size),)
    # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)
    smem_layout = cute.logical_divide(smem_layout, tiler_inst)

    # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE)
    sfa_smem_layout_staged = cute.append(
        smem_layout,
        cute.make_layout(
            num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
        ),
    )

    return sfa_smem_layout_staged


def make_smem_layout_sfb_custom(
    tiled_mma: cute.TiledMma,
    mma_tiler_mnk: Tuple[int, int, int],
    sf_vec_size: int,
    num_stages: int,
    mma_inst_tile_k: int = 4,  # NEW: Configurable parameter
) -> cute.Layout:
    """
    Custom version of blockscaled_utils.make_smem_layout_sfb that accepts
    mma_inst_tile_k as a parameter instead of hardcoding it to 4.
    """
    # (Round_Up(CTA_Tile_Shape_N, 128), MMA_Tile_Shape_K)
    sfb_tile_shape = (
        cute.round_up(mma_tiler_mnk[1], 128),
        mma_tiler_mnk[2],
    )

    # ((Atom_N, Rest_N),(Atom_K, Rest_K))
    smem_layout = cute.tile_to_shape(
        BlockScaledBasicChunk(sf_vec_size).layout,
        sfb_tile_shape,
        (2, 1),
    )

    # Use the passed mma_inst_tile_k instead of hardcoded 4
    # (CTA_Tile_Shape_N, MMA_Inst_Shape_K)
    sfb_tile_shape = cute.shape_div(sfb_tile_shape, (1, mma_inst_tile_k))
    # ((Atom_Inst_N, Atom_Inst_K), MMA_N, MMA_K)
    smem_layout = cute.tiled_divide(smem_layout, sfb_tile_shape)

    atom_n = 128
    tiler_inst = ((atom_n, sf_vec_size),)
    # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)
    smem_layout = cute.logical_divide(smem_layout, tiler_inst)

    # (((Atom_Inst_N, Rest_N),(Atom_Inst_K, Rest_K)), MMA_N, MMA_K, STAGE)
    sfb_smem_layout_staged = cute.append(
        smem_layout,
        cute.make_layout(
            num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
        ),
    )

    return sfb_smem_layout_staged


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

    V10 Architecture:
    - Independent work loops per warp role (TMA, MMA, Epilogue)
    - No global synchronization barrier between work items
    - StaticPersistentTileScheduler per warp role
    - TMEM allocation only for MMA and Epilogue warps
    """

    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,
        mma_inst_tile_k: int = 4,  # NEW: Configurable K-tile multiplier
        ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN,
        sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN,
        c_dtype: Type[cutlass.Numeric] = cutlass.Float16,
    ):
        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.mma_inst_tile_k = mma_inst_tile_k  # Store for use in layout functions
        self.num_ab_stage_override = num_ab_stage
        self.num_c_stage_override = num_c_stage

        # Store data types as instance variables
        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
        # Epilogue sync barrier - only epilogue warps participate
        self.epilog_sync_barrier = pipeline.NamedBarrier(
            barrier_id=1,
            num_threads=32 * len(self.epilog_warp_id),
        )
        # TMEM allocation barrier - CRITICAL: Only MMA + Epilogue warps (excludes TMA warp)
        self.tmem_alloc_barrier = pipeline.NamedBarrier(
            barrier_id=2,
            num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
        )

        # SM100 has 228KB shared memory
        self.smem_capacity = 228 * 1024
        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])
        # Use configurable mma_inst_tile_k from __init__
        self.mma_tiler = (
            self.mma_inst_shape_mn[0],
            self.mma_inst_shape_mn[1],
            mma_inst_shape_k * self.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 * self.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(30, 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,
        )
        # Use custom SMEM layout functions that accept mma_inst_tile_k
        self.sfa_smem_layout_staged = make_smem_layout_sfa_custom(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
            self.mma_inst_tile_k,
        )
        self.sfb_smem_layout_staged = make_smem_layout_sfb_custom(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
            self.mma_inst_tile_k,
        )
        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
        # NOTE: TMEM columns must always be computed with factor 4, NOT mma_inst_tile_k.
        # The SMEM layout shape is constant (due to divide-by-mma_inst_tile_k cancellation),
        # so TMEM layout derived from SMEM always expects the same number of columns.
        # Using mma_inst_tile_k here causes buffer mismatch when mma_inst_tile_k != 4.
        sf_atom_mn = 32
        tmem_k_factor = 4  # Hardware constant, not mma_inst_tile_k
        # Double-buffer scale factors in TMEM for S2T pipelining
        self.num_sf_stages = 2
        self.num_sfa_tmem_cols_per_stage = (
            self.cta_tile_shape_mnk[0] // sf_atom_mn
        ) * tmem_k_factor
        self.num_sfb_tmem_cols_per_stage = (
            self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
        ) * tmem_k_factor
        # Total SF columns = (SFA + SFB) * num_stages
        self.num_sf_tmem_cols = (
            self.num_sfa_tmem_cols_per_stage + self.num_sfb_tmem_cols_per_stage
        ) * self.num_sf_stages
        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)
        # Use custom layout functions with configurable mma_inst_tile_k
        sfa_smem_layout = make_smem_layout_sfa_custom(
            tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k
        )
        sfb_smem_layout = make_smem_layout_sfb_custom(
            tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k
        )
        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

        # Use overlapping accumulator only for wide N tiles (matches reference kernels).
        num_acc_stage = 1 if mma_tiler[1] == 256 else 2

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

        # Fixed overhead for barriers
        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(30, available // ab_sf_bytes_per_stage))

        return num_acc_stage, num_ab_stage, num_c_stage

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

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

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

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

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

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

        # Setup kernel attributes
        self._setup_attributes()

        # Setup scale factor tensor layouts
        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,
        )

        # Compute grid using PersistentTileScheduler
        c_shape = cute.slice_(self.cta_tile_shape_mnk, (None, None, 0))
        gc = cute.zipped_divide(c_tensor, tiler=c_shape)
        num_ctas_mnl = gc[(0, (None, None, None))].shape
        cluster_shape_mnk = (*self.cluster_shape_mn, 1)

        tile_sched_params = utils.PersistentTileSchedulerParams(
            num_ctas_mnl, cluster_shape_mnk
        )
        grid = utils.StaticPersistentTileScheduler.get_grid_shape(
            tile_sched_params, max_active_clusters
        )

        self.buffer_align_bytes = 1024

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

        self.shared_storage = SharedStorage

        self.kernel(
            tiled_mma,
            tiled_mma_sfb,
            tma_atom_a,
            tma_tensor_a,
            tma_atom_b,
            tma_tensor_b,
            tma_atom_sfa,
            tma_tensor_sfa,
            tma_atom_sfb,
            tma_tensor_sfb,
            tma_atom_c,
            tma_tensor_c,
            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,
            tile_sched_params,
            epilogue_op,
        ).launch(
            grid=grid,
            block=[self.threads_per_cta, 1, 1],
            cluster=(*self.cluster_shape_mn, 1),
            min_blocks_per_mp=1,
        )

    @cute.kernel
    def kernel(
        self,
        tiled_mma: cute.TiledMma,
        tiled_mma_sfb: cute.TiledMma,
        tma_atom_a: cute.CopyAtom,
        mA_mkl: cute.Tensor,
        tma_atom_b: cute.CopyAtom,
        mB_nkl: cute.Tensor,
        tma_atom_sfa: cute.CopyAtom,
        mSFA_mkl: cute.Tensor,
        tma_atom_sfb: cute.CopyAtom,
        mSFB_nkl: cute.Tensor,
        tma_atom_c: cute.CopyAtom,
        mC_mnl: cute.Tensor,
        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,
        tile_sched_params: utils.PersistentTileSchedulerParams,
        epilogue_op: cutlass.Constexpr,
    ):
        """
        GPU device kernel for block-scaled GEMM with independent work loops per warp role.
        """
        warp_idx = cute.arch.warp_idx()
        warp_idx = cute.arch.make_warp_uniform(warp_idx)

        # TMA descriptor prefetch
        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
        bidx, bidy, bidz = cute.arch.block_idx()
        mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
        is_leader_cta = mma_tile_coord_v == 0
        cta_rank_in_cluster = cute.arch.make_warp_uniform(
            cute.arch.block_idx_in_cluster()
        )
        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()

        # Shared memory allocation
        smem = utils.SmemAllocator()
        storage = smem.allocate(self.shared_storage)

        # Initialize mainloop ab_pipeline
        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,
        )

        # Initialize acc_pipeline
        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
        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
        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
        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)
        )
        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 for A
        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)

        # =================================================================
        # INDEPENDENT WORK LOOPS PER WARP ROLE
        # No global synchronization barrier between work items
        # =================================================================

        # -----------------------------------------------------------------
        # TMA LOAD WARP (warp 5) - INDEPENDENT WORK LOOP
        # -----------------------------------------------------------------
        if warp_idx == self.tma_warp_id:
            # Create tile scheduler for TMA warp
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )
            work_tile = tile_sched.initial_work_tile_info()

            ab_producer_state = pipeline.make_pipeline_state(
                pipeline.PipelineUserType.Producer, self.num_ab_stage
            )

            while work_tile.is_valid_tile:
                # Get tile coord from tile scheduler
                cur_tile_coord = work_tile.tile_idx
                mma_tile_coord_mnl = (
                    cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                    cur_tile_coord[1],
                    cur_tile_coord[2],
                )

                # Slice to per mma tile index
                tAgA_slice = tAgA[
                    (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                ]
                tBgB_slice = tBgB[
                    (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
                ]
                tAgSFA_slice = tAgSFA[
                    (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                ]

                slice_n = mma_tile_coord_mnl[1]
                if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
                    slice_n = mma_tile_coord_mnl[1] // 2
                tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]

                # Reset and peek for first k tile
                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
                    )

                # TMA load loop
                for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
                    # Wait for AB buffer empty
                    ab_pipeline.producer_acquire(
                        ab_producer_state, peek_ab_empty_status
                    )

                    # TMA loads
                    cute.copy(
                        tma_atom_a,
                        tAgA_slice[(None, ab_producer_state.count)],
                        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, ab_producer_state.count)],
                        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, ab_producer_state.count)],
                        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, ab_producer_state.count)],
                        tBsSFB[(None, ab_producer_state.index)],
                        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
                        mcast_mask=sfb_full_mcast_mask,
                    )

                    # Advance and peek for next k tile
                    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
                        )

                # Advance to next tile
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()

            # TMA warp tail
            ab_pipeline.producer_tail(ab_producer_state)

        # -----------------------------------------------------------------
        # MMA WARP (warp 4) - INDEPENDENT WORK LOOP
        # -----------------------------------------------------------------
        if warp_idx == self.mma_warp_id:
            # Wait for TMEM allocation
            tmem.wait_for_alloc()

            # Retrieve TMEM pointers
            acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

            # Create TMEM layouts for scale factors
            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)),
            )
            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)),
            )

            # Double-buffered SFA TMEM tensors (buffer 0 and buffer 1)
            # Layout: [Acc] [SFA_0] [SFA_1] [SFB_0] [SFB_1]
            # Default to manual offsets to keep overlap layout invariants intact.
            sfa_tmem_ptr_0 = cute.recast_ptr(
                acc_tmem_ptr + self.num_accumulator_tmem_cols,
                dtype=self.sf_dtype,
            )
            sfa_tmem_ptr_1 = cute.recast_ptr(
                acc_tmem_ptr
                + self.num_accumulator_tmem_cols
                + self.num_sfa_tmem_cols_per_stage,
                dtype=self.sf_dtype,
            )
            tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
            tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)

            sfb_base_offset = (
                self.num_accumulator_tmem_cols
                + self.num_sfa_tmem_cols_per_stage * self.num_sf_stages
            )
            sfb_tmem_ptr_0 = cute.recast_ptr(
                acc_tmem_ptr + sfb_base_offset,
                dtype=self.sf_dtype,
            )
            sfb_tmem_ptr_1 = cute.recast_ptr(
                acc_tmem_ptr
                + sfb_base_offset
                + self.num_sfb_tmem_cols_per_stage,
                dtype=self.sf_dtype,
            )
            tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
            tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)
            sfb_tmem_cols = self.num_sfb_tmem_cols_per_stage

            # Non-overlapping path can safely use layout-derived offsets.
            if cutlass.const_expr(not self.overlapping_accum):
                acc_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
                sfa_tmem_ptr_0 = cute.recast_ptr(
                    acc_tmem_ptr + acc_tmem_cols,
                    dtype=self.sf_dtype,
                )
                tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
                sfa_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFA_0)
                sfa_tmem_ptr_1 = cute.recast_ptr(
                    acc_tmem_ptr + acc_tmem_cols + sfa_tmem_cols,
                    dtype=self.sf_dtype,
                )
                tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)

                sfb_base_offset = acc_tmem_cols + sfa_tmem_cols * self.num_sf_stages
                sfb_tmem_ptr_0 = cute.recast_ptr(
                    acc_tmem_ptr + sfb_base_offset,
                    dtype=self.sf_dtype,
                )
                tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
                sfb_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFB_0)
                sfb_tmem_ptr_1 = cute.recast_ptr(
                    acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols,
                    dtype=self.sf_dtype,
                )
                tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)

            # S2T copy partitions for scale factors (both buffers)
            (
                tiled_copy_s2t_sfa,
                tCsSFA_compact_s2t,
                tCtSFA_compact_s2t_0,
            ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_0)
            (
                _,
                _,
                tCtSFA_compact_s2t_1,
            ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_1)
            (
                tiled_copy_s2t_sfb,
                tCsSFB_compact_s2t,
                tCtSFB_compact_s2t_0,
            ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_0)
            (
                _,
                _,
                tCtSFB_compact_s2t_1,
            ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_1)

            # Create tile scheduler for MMA warp
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )
            work_tile = tile_sched.initial_work_tile_info()

            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
            )

            while work_tile.is_valid_tile:
                # Get tile coord from tile scheduler
                cur_tile_coord = work_tile.tile_idx
                mma_tile_coord_mnl = (
                    cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                    cur_tile_coord[1],
                    cur_tile_coord[2],
                )

                # Get accumulator stage index
                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)]

                # Peek for first k tile
                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
                    )

                # Wait for accumulator buffer empty
                if is_leader_cta:
                    acc_pipeline.producer_acquire(acc_producer_state)

                # Handle SFB offset for different tile shapes (double-buffered)
                # Create SFB MMA tensors for both buffers with appropriate offsets
                tCtSFB_mma_0 = tCtSFB_0
                tCtSFB_mma_1 = tCtSFB_1
                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_0 = cute.recast_ptr(
                        acc_tmem_ptr + sfb_base_offset + offset,
                        dtype=self.sf_dtype,
                    )
                    shifted_ptr_1 = cute.recast_ptr(
                        acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
                        dtype=self.sf_dtype,
                    )
                    tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
                    tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, 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_0 = cute.recast_ptr(
                        acc_tmem_ptr + sfb_base_offset + offset,
                        dtype=self.sf_dtype,
                    )
                    shifted_ptr_1 = cute.recast_ptr(
                        acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
                        dtype=self.sf_dtype,
                    )
                    tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
                    tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, tCtSFB_layout)

                # Reset accumulate flag
                tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

                # MMA mainloop with double-buffered SF TMEM
                for k_tile in range(k_tile_cnt):
                    # Default to buffer 0 so variables exist outside dynamic control flow.
                    tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
                    tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
                    tCtSFA_mma = tCtSFA_0
                    tCtSFB_mma = tCtSFB_mma_0

                    if is_leader_cta:
                        # Select SF TMEM buffer based on pipeline stage parity
                        sf_stage_is_zero = (ab_consumer_state.index & 1) == 0
                        if sf_stage_is_zero:
                            tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
                            tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
                            tCtSFA_mma = tCtSFA_0
                            tCtSFB_mma = tCtSFB_mma_0
                        else:
                            tCtSFA_compact_s2t = tCtSFA_compact_s2t_1
                            tCtSFB_compact_s2t = tCtSFB_compact_s2t_1
                            tCtSFA_mma = tCtSFA_1
                            tCtSFB_mma = tCtSFB_mma_1

                        # Wait for AB buffer full
                        ab_pipeline.consumer_wait(
                            ab_consumer_state, peek_ab_full_status
                        )

                        # Copy SFA/SFB from smem to tmem (using buffer 0)
                        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,
                        )

                        # MMA over k blocks
                        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_mma[sf_kblock_coord].iterator
                            )
                            tiled_mma.set(
                                tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator
                            )

                            cute.gemm(
                                tiled_mma,
                                tCtAcc,
                                tCrA[kblock_coord],
                                tCrB[kblock_coord],
                                tCtAcc,
                            )

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

                        # Release AB buffer
                        ab_pipeline.consumer_release(ab_consumer_state)

                    # Peek for next k tile
                    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
                            )

                # Commit accumulator
                if is_leader_cta:
                    acc_pipeline.producer_commit(acc_producer_state)
                acc_producer_state.advance()

                # Advance to next tile
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()

            # MMA warp tail
            acc_pipeline.producer_tail(acc_producer_state)

        # -----------------------------------------------------------------
        # EPILOGUE WARPS (warps 0-3) - INDEPENDENT WORK LOOP
        # -----------------------------------------------------------------
        if warp_idx < self.mma_warp_id:
            # Allocate TMEM
            tmem.allocate(self.num_tmem_alloc_cols)

            # Wait for TMEM allocation
            tmem.wait_for_alloc()

            # Retrieve TMEM pointer
            acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

            # Partition for epilogue
            epi_tidx = tidx
            (
                tiled_copy_t2r,
                tTR_tAcc_base,
                tTR_rAcc,
            ) = self.epilog_tmem_copy_and_partition(
                epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
            )

            tTR_rC = cute.make_rmem_tensor(tTR_rAcc.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)

            # Create tile scheduler for Epilogue warps
            tile_sched = utils.StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )
            work_tile = tile_sched.initial_work_tile_info()

            acc_consumer_state = pipeline.make_pipeline_state(
                pipeline.PipelineUserType.Consumer, self.num_acc_stage
            )

            # C TMA store pipeline
            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,
            )

            while work_tile.is_valid_tile:
                # Get tile coord from tile scheduler
                cur_tile_coord = work_tile.tile_idx
                mma_tile_coord_mnl = (
                    cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                    cur_tile_coord[1],
                    cur_tile_coord[2],
                )

                # Slice output tensor
                bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]

                # Get accumulator stage index
                if cutlass.const_expr(self.overlapping_accum):
                    acc_stage_index = acc_consumer_state.phase
                    reverse_subtile = (
                        cutlass.Boolean(True)
                        if acc_stage_index == 0
                        else cutlass.Boolean(False)
                    )
                else:
                    acc_stage_index = acc_consumer_state.index

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

                # Wait for accumulator buffer full
                acc_pipeline.consumer_wait(acc_consumer_state)

                tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
                bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))

                # Store accumulator to global memory in subtiles
                subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
                num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt

                for subtile_idx in cutlass.range(subtile_cnt):
                    real_subtile_idx = subtile_idx
                    if cutlass.const_expr(self.overlapping_accum):
                        if reverse_subtile:
                            real_subtile_idx = (
                                self.cta_tile_shape_mnk[1] // self.epi_tile_n
                                - 1
                                - subtile_idx
                            )

                    # Load accumulator from TMEM to registers
                    tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
                    cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

                    # Early release for overlapping accum
                    if cutlass.const_expr(self.overlapping_accum):
                        if subtile_idx == self.iter_acc_early_release_in_epilogue:
                            cute.arch.fence_view_async_tmem_load()
                            with cute.arch.elect_one():
                                acc_pipeline.consumer_release(acc_consumer_state)
                            acc_consumer_state.advance()

                    # Convert to C type
                    acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
                    acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
                    tRS_rC.store(acc_vec)

                    # Store C to shared memory
                    c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
                    cute.copy(
                        tiled_copy_r2s,
                        tRS_rC,
                        tRS_sC[(None, None, None, c_buffer)],
                    )

                    # Fence and barrier
                    cute.arch.fence_proxy(
                        cute.arch.ProxyKind.async_shared,
                        space=cute.arch.SharedSpace.shared_cta,
                    )
                    self.epilog_sync_barrier.arrive_and_wait()

                    # TMA store C
                    if warp_idx == self.epilog_warp_id[0]:
                        cute.copy(
                            tma_atom_c,
                            bSG_sC[(None, c_buffer)],
                            bSG_gC[(None, real_subtile_idx)],
                        )
                        c_pipeline.producer_commit()
                        c_pipeline.producer_acquire()
                    self.epilog_sync_barrier.arrive_and_wait()

                # Release accumulator buffer (non-overlapping case)
                if cutlass.const_expr(not self.overlapping_accum):
                    with cute.arch.elect_one():
                        acc_pipeline.consumer_release(acc_consumer_state)
                    acc_consumer_state.advance()

                # Advance to next tile
                tile_sched.advance_to_next_work()
                work_tile = tile_sched.get_current_work()

            # Epilogue tail - dealloc TMEM and wait for C stores
            tmem.relinquish_alloc_permit()
            self.epilog_sync_barrier.arrive_and_wait()
            tmem.free(acc_tmem_ptr)
            c_pipeline.producer_tail()

    def mainloop_s2t_copy_and_partition(self, sSF, tCtSF):
        """Create S2T copy for scale factors."""
        tCsSF_compact = cute.filter_zeros(sSF)
        tCtSF_compact = cute.filter_zeros(tCtSF)

        copy_atom_s2t = cute.make_copy_atom(
            tcgen05.Cp4x32x128bOp(self.cta_group),
            self.sf_dtype,
        )
        tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
        thr_copy_s2t = tiled_copy_s2t.get_slice(0)

        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_
        )
        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."""
        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,
        )
        tAcc_epi = cute.flat_divide(
            tAcc[((None, None), 0, 0, None)],
            epi_tile,
        )
        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)
        tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)

        gC_mnl_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        tTR_rAcc = 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_rAcc

    def epilog_smem_copy_and_partition(self, tiled_copy_t2r, tTR_rC, tidx, sC):
        """Partition for register to SMEM copy in epilogue."""
        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)
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        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."""
        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)
        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


# =============================================================================
# Kernel Compilation and Dispatch
# =============================================================================

_compiled_kernel_cache = {}

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


@dataclass(frozen=True)
class KernelConfig:
    mma_tiler_mn: Tuple[int, int]
    cluster_shape_mn: Tuple[int, int]
    num_ab_stage: int = 0  # 0 => auto (use _compute_stages())
    num_c_stage: int = 0  # 0 => auto (use _compute_stages())
    mma_inst_tile_k: int = 4  # K-tile multiplier: 1, 2, or 4


_OPTIMAL_CONFIG = {
    (128, 7168, 16384, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
    ),
    (128, 4096, 7168, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
    ),
    (128, 7168, 2048, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
    ),
}

# Default fallback configuration for unseen shapes.
# mma_inst_tile_k=2 gives smaller K-tiles (128 vs 256), allowing more AB stages
_DEFAULT_CONFIG = KernelConfig(
    (128, 64),
    (1, 1),
    num_ab_stage=10,
    num_c_stage=2,
    mma_inst_tile_k=2,
)


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


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


def _ensure_compiled(
    m: int,
    n: int,
    k: int,
    l: int,
    mma_tiler_mn: Tuple[int, int],
    cluster_shape_mn: Tuple[int, int],
    num_ab_stage: int,
    num_c_stage: int,
    mma_inst_tile_k: int,
) -> None:
    """Compile a single (shape, tiling) variant if it's missing."""
    global _compiled_kernel_cache
    key = _cache_key(
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        num_ab_stage,
        num_c_stage,
        mma_inst_tile_k,
    )
    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,
        mma_inst_tile_k=mma_inst_tile_k,
        ab_dtype=ab_dtype,
        sf_dtype=sf_dtype,
        c_dtype=c_dtype,
    )

    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)

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


def compile_kernel():
    """Pre-compile ALL required kernel configurations."""
    global _compiled_kernel_cache

    if _compiled_kernel_cache is None:
        _compiled_kernel_cache = {}

    for shape_key, cfg in _OPTIMAL_CONFIG.items():
        m, n, k, l = shape_key
        num_ab_stage = max(0, int(cfg.num_ab_stage))
        num_c_stage = max(0, int(cfg.num_c_stage))
        mma_inst_tile_k = int(cfg.mma_inst_tile_k)
        _ensure_compiled(
            m,
            n,
            k,
            l,
            cfg.mma_tiler_mn,
            cfg.cluster_shape_mn,
            num_ab_stage,
            num_c_stage,
            mma_inst_tile_k,
        )


_gmem = cute.AddressSpace.gmem


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

    global _compiled_kernel_cache
    if _compiled_kernel_cache is None:
        _compiled_kernel_cache = {}

    m, n, l = c.shape
    k = a.shape[1] << 1

    shape_key = (m, n, k, l)
    cfg = _OPTIMAL_CONFIG.get(shape_key, _DEFAULT_CONFIG)
    mma_tiler_mn = cfg.mma_tiler_mn
    cluster_shape_mn = cfg.cluster_shape_mn
    num_ab_stage = _env_override_int("AB_STAGES")
    num_ab_stage = int(cfg.num_ab_stage) if num_ab_stage is None else num_ab_stage
    if num_ab_stage < 0:
        num_ab_stage = 0
    num_c_stage = _env_override_int("C_STAGES")
    num_c_stage = int(cfg.num_c_stage) if num_c_stage is None else num_c_stage
    if num_c_stage < 0:
        num_c_stage = 0
    mma_inst_tile_k = int(cfg.mma_inst_tile_k)

    cache_key = _cache_key(
        m,
        n,
        k,
        l,
        mma_tiler_mn,
        cluster_shape_mn,
        num_ab_stage,
        num_c_stage,
        mma_inst_tile_k,
    )

    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)

    if _c[6] != cache_key or _c[0] is None:
        if cache_key not in _compiled_kernel_cache:
            _ensure_compiled(
                m,
                n,
                k,
                l,
                mma_tiler_mn,
                cluster_shape_mn,
                num_ab_stage,
                num_c_stage,
                mma_inst_tile_k,
            )
        _c[0] = _compiled_kernel_cache[cache_key]
        _c[6] = cache_key

    _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

    _c[0](
        _c[1],
        _c[2],
        _c[3],
        _c[4],
        _c[5],
        (m, n, k, l),
    )
    return c
scrolls · 1793 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 168494.

⋯ 352 unchanged lines
self.overlapping_accum = self.num_acc_stage == 1
# TMEM column counts
+ # NOTE: TMEM columns must always be computed with factor 4, NOT mma_inst_tile_k.
+ # The SMEM layout shape is constant (due to divide-by-mma_inst_tile_k cancellation),
+ # so TMEM layout derived from SMEM always expects the same number of columns.
+ # Using mma_inst_tile_k here causes buffer mismatch when mma_inst_tile_k != 4.
sf_atom_mn = 32
- self.num_sfa_tmem_cols = (
+ tmem_k_factor = 4 # Hardware constant, not mma_inst_tile_k
+ # Double-buffer scale factors in TMEM for S2T pipelining
+ self.num_sf_stages = 2
+ self.num_sfa_tmem_cols_per_stage = (
self.cta_tile_shape_mnk[0] // sf_atom_mn
- ) * self.mma_inst_tile_k
- self.num_sfb_tmem_cols = (
+ ) * tmem_k_factor
+ self.num_sfb_tmem_cols_per_stage = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
- ) * self.mma_inst_tile_k
- self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
+ ) * tmem_k_factor
+ # Total SF columns = (SFA + SFB) * num_stages
+ self.num_sf_tmem_cols = (
+ self.num_sfa_tmem_cols_per_stage + self.num_sfb_tmem_cols_per_stage
+ ) * self.num_sf_stages
self.num_accumulator_tmem_cols = (
self.cta_tile_shape_mnk[1] * self.num_acc_stage
if not self.overlapping_accum
⋯ 38 unchanged lines
ab_sf_bytes_per_stage = a_bytes + b_bytes + sfa_bytes + sfb_bytes
- # Use num_acc_stage=1 for overlapping accumulator
- num_acc_stage = 1
+ # Use overlapping accumulator only for wide N tiles (matches reference kernels).
+ num_acc_stage = 1 if mma_tiler[1] == 256 else 2
# Compute C stages (2 for double buffering)
num_c_stage = 2
⋯ 605 unchanged lines
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
- # Make SFA tmem tensor
- sfa_tmem_ptr = cute.recast_ptr(
- acc_tmem_ptr + self.num_accumulator_tmem_cols,
- dtype=self.sf_dtype,
- )
+ # Create TMEM layouts for scale factors
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)
-
- # Make SFB tmem tensor
- 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)
- # S2T copy partitions for scale factors
+ # Double-buffered SFA TMEM tensors (buffer 0 and buffer 1)
+ # Layout: [Acc] [SFA_0] [SFA_1] [SFB_0] [SFB_1]
+ # Default to manual offsets to keep overlap layout invariants intact.
+ sfa_tmem_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + self.num_accumulator_tmem_cols,
+ dtype=self.sf_dtype,
+ )
+ sfa_tmem_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr
+ + self.num_accumulator_tmem_cols
+ + self.num_sfa_tmem_cols_per_stage,
+ dtype=self.sf_dtype,
+ )
+ tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
+ tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)
+
+ sfb_base_offset = (
+ self.num_accumulator_tmem_cols
+ + self.num_sfa_tmem_cols_per_stage * self.num_sf_stages
+ )
+ sfb_tmem_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset,
+ dtype=self.sf_dtype,
+ )
+ sfb_tmem_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr
+ + sfb_base_offset
+ + self.num_sfb_tmem_cols_per_stage,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
+ tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)
+ sfb_tmem_cols = self.num_sfb_tmem_cols_per_stage
+
+ # Non-overlapping path can safely use layout-derived offsets.
+ if cutlass.const_expr(not self.overlapping_accum):
+ acc_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ sfa_tmem_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + acc_tmem_cols,
+ dtype=self.sf_dtype,
+ )
+ tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
+ sfa_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFA_0)
+ sfa_tmem_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr + acc_tmem_cols + sfa_tmem_cols,
+ dtype=self.sf_dtype,
+ )
+ tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)
+
+ sfb_base_offset = acc_tmem_cols + sfa_tmem_cols * self.num_sf_stages
+ sfb_tmem_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
+ sfb_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFB_0)
+ sfb_tmem_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)
+
+ # S2T copy partitions for scale factors (both buffers)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
- tCtSFA_compact_s2t,
- ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
+ tCtSFA_compact_s2t_0,
+ ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_0)
(
+ _,
+ _,
+ tCtSFA_compact_s2t_1,
+ ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_1)
+ (
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
- tCtSFB_compact_s2t,
- ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
+ tCtSFB_compact_s2t_0,
+ ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_0)
+ (
+ _,
+ _,
+ tCtSFB_compact_s2t_1,
+ ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_1)
# Create tile scheduler for MMA warp
tile_sched = utils.StaticPersistentTileScheduler.create(
⋯ 37 unchanged lines
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
- # Handle SFB offset for different tile shapes
- tCtSFB_mma = tCtSFB
+ # Handle SFB offset for different tile shapes (double-buffered)
+ # Create SFB MMA tensors for both buffers with appropriate offsets
+ tCtSFB_mma_0 = tCtSFB_0
+ tCtSFB_mma_1 = tCtSFB_1
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,
+ shifted_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset + offset,
dtype=self.sf_dtype,
)
- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
+ shifted_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
+ tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, 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,
+ shifted_ptr_0 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset + offset,
dtype=self.sf_dtype,
)
- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
+ shifted_ptr_1 = cute.recast_ptr(
+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
+ tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, tCtSFB_layout)
# Reset accumulate flag
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
- # MMA mainloop
+ # MMA mainloop with double-buffered SF TMEM
for k_tile in range(k_tile_cnt):
+ # Default to buffer 0 so variables exist outside dynamic control flow.
+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
+ tCtSFA_mma = tCtSFA_0
+ tCtSFB_mma = tCtSFB_mma_0
+
if is_leader_cta:
+ # Select SF TMEM buffer based on pipeline stage parity
+ sf_stage_is_zero = (ab_consumer_state.index & 1) == 0
+ if sf_stage_is_zero:
+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
+ tCtSFA_mma = tCtSFA_0
+ tCtSFB_mma = tCtSFB_mma_0
+ else:
+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_1
+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_1
+ tCtSFA_mma = tCtSFA_1
+ tCtSFB_mma = tCtSFB_mma_1
+
# Wait for AB buffer full
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
- # Copy SFA/SFB from smem to tmem
+ # Copy SFA/SFB from smem to tmem (using buffer 0)
s2t_stage_coord = (
None,
None,
⋯ 27 unchanged lines
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(
- tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator
+ tcgen05.Field.SFA, tCtSFA_mma[sf_kblock_coord].iterator
)
tiled_mma.set(
tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator
⋯ 299 unchanged lines
_OPTIMAL_CONFIG = {
(128, 7168, 16384, 1): KernelConfig(
- (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
),
(128, 4096, 7168, 1): KernelConfig(
- (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
),
(128, 7168, 2048, 1): KernelConfig(
- (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
),
}
# Default fallback configuration for unseen shapes.
+ # mma_inst_tile_k=2 gives smaller K-tiles (128 vs 256), allowing more AB stages
_DEFAULT_CONFIG = KernelConfig(
(128, 64),
(1, 1),
- num_ab_stage=7,
+ num_ab_stage=10,
num_c_stage=2,
- mma_inst_tile_k=4,
+ mma_inst_tile_k=2,
)
scrolls · 289 diff lines total

Best evidence level for this revision: reported

JSON