Skip to content
KernelIndex
Search⌘K

submission 168494

rcmalli · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-168494?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
#86 of 369
2025-12-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8f31d8dbf2391801aa1ce1ef9c98ed9f91fd3f755781f4eb458d57a710cd3348
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.py1697 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
        sf_atom_mn = 32
        self.num_sfa_tmem_cols = (
            self.cta_tile_shape_mnk[0] // sf_atom_mn
        ) * self.mma_inst_tile_k
        self.num_sfb_tmem_cols = (
            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
        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 num_acc_stage=1 for overlapping accumulator
        num_acc_stage = 1

        # 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)

            # Make SFA tmem tensor
            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)

            # 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
            (
                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)

            # 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
                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)

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

                # MMA mainloop
                for k_tile in range(k_tile_cnt):
                    if is_leader_cta:
                        # Wait for AB buffer full
                        ab_pipeline.consumer_wait(
                            ab_consumer_state, peek_ab_full_status
                        )

                        # Copy SFA/SFB from smem to tmem
                        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[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=2, 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, 7168, 2048, 1): KernelConfig(
        (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
    ),
}

# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
    (128, 64),
    (1, 1),
    num_ab_stage=7,
    num_c_stage=2,
    mma_inst_tile_k=4,
)


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 · 1697 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 155565.

⋯ 3 unchanged lines
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
+ 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
⋯ 14 unchanged lines
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 SchedulerParams:
- tiles_m: int
- tiles_n: int
- k_tile_cnt: int
+ 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 tiles_mn(self) -> int:
- return self.tiles_m * self.tiles_n
+ 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 total_work_items(self, l: int) -> int:
- return self.tiles_mn * l
- @classmethod
- def from_shape(
- cls,
- m: int,
- n: int,
- k_tile_cnt: int,
- tile_m: int,
- tile_n: int,
- ):
- tiles_m = (m + tile_m - 1) // tile_m
- tiles_n = (n + tile_n - 1) // tile_n
- return cls(tiles_m, tiles_n, k_tile_cnt)
+ 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).
- 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)
+ 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__(
⋯ 3 unchanged lines
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,
- 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.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 (needed for JIT compilation)
+ # Store data types as instance variables
self.a_dtype = ab_dtype
self.b_dtype = ab_dtype
self.sf_dtype = sf_dtype
⋯ 14 unchanged lines
)
# 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=self.threads_per_cta,
+ num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
)
- # 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 has 228KB shared memory
+ self.smem_capacity = 228 * 1024
SM100_TMEM_CAPACITY_COLUMNS = 512
self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS
⋯ 29 unchanged lines
# Compute MMA/cluster/tile shapes
mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
- mma_inst_tile_k = 4
+ # 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 * mma_inst_tile_k,
+ 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 * mma_inst_tile_k,
+ 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),
⋯ 48 unchanged lines
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)))
+ 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)))
⋯ 10 unchanged lines
self.b_dtype,
self.num_ab_stage,
)
- self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
+ # 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 = blockscaled_utils.make_smem_layout_sfb(
+ 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,
⋯ 9 unchanged lines
sf_atom_mn = 32
self.num_sfa_tmem_cols = (
self.cta_tile_shape_mnk[0] // sf_atom_mn
- ) * mma_inst_tile_k
+ ) * self.mma_inst_tile_k
self.num_sfb_tmem_cols = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
- ) * mma_inst_tile_k
+ ) * self.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
⋯ 22 unchanged lines
# 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
+ # 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 = blockscaled_utils.make_smem_layout_sfb(
- tiled_mma, mma_tiler, sf_vec_size, 1
+ 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)
⋯ 5 unchanged lines
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
+ # Use num_acc_stage=1 for overlapping accumulator
num_acc_stage = 1
# Compute C stages (2 for double buffering)
num_c_stage = 2
- # Fixed overhead for barriers, etc.
+ # 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(8, available // ab_sf_bytes_per_stage))
+ num_ab_stage = max(2, min(30, 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,
⋯ 6 unchanged lines
max_active_clusters: cutlass.Constexpr = 132,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
- """Execute the block-scaled GEMM operation.
-
- Args:
- a_ptr: Pointer to A matrix data (M x K x L, K-major, FP4)
- b_ptr: Pointer to B matrix data (N x K x L, K-major, FP4)
- sfa_ptr: Pointer to scale factors for A (permuted layout, FP8)
- sfb_ptr: Pointer to scale factors for B (permuted layout, FP8)
- c_ptr: Pointer to C matrix data (M x N x L, N-major, FP16)
- shape_mnkl: Tuple of (M, N, K, L) dimensions
- max_active_clusters: Maximum active clusters for persistent scheduling
- epilogue_op: Optional epilogue operation
- """
+ """Execute the block-scaled GEMM operation."""
m, n, k, l = shape_mnkl
# Create tensors from pointers with proper layouts
⋯ 5 unchanged lines
b_layout = cute.make_layout((n, k, l), stride=(k, 1, n * k))
b_tensor = cute.make_tensor(b_ptr, b_layout)
- # C is (M, N, L) row-major (N-stride=1, M-stride=N)
+ # C 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)
⋯ 9 unchanged lines
# 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
+ # Setup scale factor tensor layouts
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
a_tensor.shape, self.sf_vec_size
)
⋯ 102 unchanged lines
self.epi_tile,
)
- # ---------------------------------------------------------------------
- # Deterministic persistent scheduling
- #
- # Launch a single logical cluster in (x,y) and use grid.z as a persistent
- # work queue length. All warps participate in a shared work loop and
- # synchronize once per work item.
- # ---------------------------------------------------------------------
- cluster_m, cluster_n = self.cluster_shape_mn
- v_ctas_per_tile = cute.size(tiled_mma.thr_id.shape)
- scheduler_params = SchedulerParams.from_shape(
- m=c_tensor.shape[0],
- n=c_tensor.shape[1],
- k_tile_cnt=k // self.mma_tiler[2],
- tile_m=self.mma_tiler[0] * cluster_m,
- tile_n=self.mma_tiler[1] * cluster_n,
+ # 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
+ )
- 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
⋯ 51 unchanged lines
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,
⋯ 2 unchanged lines
self.sfb_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
- tiles_m,
- tiles_n,
- l,
- total_work_items,
+ tile_sched_params,
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,
- ),
+ cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=1,
)
⋯ 12 unchanged lines
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,
⋯ 2 unchanged lines
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,
+ tile_sched_params: utils.PersistentTileSchedulerParams,
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.
+ GPU device kernel for block-scaled GEMM with independent work loops per warp role.
"""
- # 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.
+ # TMA descriptor prefetch
if warp_idx == self.tma_warp_id:
cpasync.prefetch_descriptor(tma_atom_a)
cpasync.prefetch_descriptor(tma_atom_b)
⋯ 3 unchanged lines
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.
+ # CTA/thread coordinates
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
+ 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()
)
- # 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
)
⋯ 1 unchanged lines
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.
+ # Shared memory allocation
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
- # Pipelines:
- # - `PipelineTmaUmma` coordinates TMA producer(s) and UMMA consumer(s) via mbarrier arrays.
+ # 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(
⋯ 9 unchanged lines
defer_sync=True,
)
- # Accumulator pipeline:
- # - `PipelineUmmaAsync` synchronizes "accumulator stage" ownership across MMA and epilogue roles.
+ # 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
⋯ 10 unchanged lines
defer_sync=True,
)
- # TMEM allocator:
- # - Provides per-CTA allocations of Blackwell tensor memory columns for accumulators + SF tiles.
+ # TMEM allocator
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=self.tmem_alloc_barrier,
⋯ 18 unchanged lines
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=...)`.
+ # TMA multicast masks
a_full_mcast_mask = None
b_full_mcast_mask = None
sfa_full_mcast_mask = None
⋯ 12 unchanged lines
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.
+ # Global -> per-CTA tiles
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
⋯ 11 unchanged lines
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
⋯ 5 unchanged lines
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.
+ # TMA partitions for A
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
⋯ 73 unchanged lines
# 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)
+ # =================================================================
+ # INDEPENDENT WORK LOOPS PER WARP ROLE
+ # No global synchronization barrier between work items
+ # =================================================================
- # 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
- )
+ # -----------------------------------------------------------------
+ # 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()
- # 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,
+ ab_producer_state = pipeline.make_pipeline_state(
+ pipeline.PipelineUserType.Producer, self.num_ab_stage
)
- # -------------------------------------------------------------
- # TMA LOAD WARP (warp 5)
- # -------------------------------------------------------------
- if warp_idx == self.tma_warp_id:
+ 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, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
+ (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB[
- (None, mma_tile_coord_mnl[1], None, cluster_tile_coord_mnl[2])
+ (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA[
- (None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
+ (None, mma_tile_coord_mnl[0], None, mma_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]
+ 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])]
+ 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:
⋯ 1 unchanged lines
ab_producer_state
)
- for k_iter in cutlass.range(0, k_tile_cnt, 1, unroll=1): # noqa: B007
+ # 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
)
- 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.
+ # TMA loads
cute.copy(
tma_atom_a,
- tAgA_slice[(None, k_tile)],
+ 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, k_tile)],
+ 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, k_tile)],
+ 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, k_tile)],
+ 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:
⋯ 1 unchanged lines
ab_producer_state
)
- # -------------------------------------------------------------
- # MMA WARP (warp 4)
- # -------------------------------------------------------------
- if warp_idx == self.mma_warp_id:
+ # 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)
+
+ # Make SFA tmem tensor
+ 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)
+
+ # 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
+ (
+ 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)
+
+ # 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:
⋯ 1 unchanged lines
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:
⋯ 1 unchanged lines
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
tCtSFB_mma = tCtSFB
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
offset = (
⋯ 20 unchanged lines
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
+ # Reset accumulate flag
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
- for k_iter in range(k_tile_cnt): # noqa: B007
+ # MMA mainloop
+ for k_tile in range(k_tile_cnt):
if is_leader_cta:
+ # Wait for AB buffer full
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
+ # Copy SFA/SFB from smem to tmem
s2t_stage_coord = (
None,
None,
⋯ 15 unchanged lines
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 = (
⋯ 11 unchanged lines
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,
⋯ 4 unchanged lines
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:
⋯ 2 unchanged lines
ab_consumer_state
)
+ # Commit accumulator
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:
+ # 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
⋯ 1 unchanged lines
(None, None, None, None, None, acc_stage_index)
]
+ # Wait for accumulator buffer full
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))
+ # Store accumulator to global memory in subtiles
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
- num_prev_subtiles = 0
+ num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
- if warp_idx == self.epilog_warp_id[0] and subtile_cnt > 0:
- c_pipeline.producer_acquire()
+ 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
+ )
- 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)
+ # 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)
- 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
+ # 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()
- 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)
+ # 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)
- 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
+ # 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)],
)
- # Memory ordering for TMA store:
- # - `cute.arch.fence_proxy(async_shared)` makes prior SMEM writes visible to async consumers (TMA).
+ # 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, subtile_idx)],
+ bSG_gC[(None, real_subtile_idx)],
⋯ diff truncated
scrolls · 1201 diff lines total

Best evidence level for this revision: reported

JSON