Skip to content
KernelIndex
Search⌘K

submission 142159

billcarson · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4a985a9db2fee19203a86a644c674f414df28914c9079f588466316984366841
license declaredunknown
license concludedunknown
authorsbillcarson
imported2026-08-15

Techniques

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

async-copydef cp_async_bulk_remote(
clusterTODO: store to dsmem in my_split_idx
fp4PyTorch reference implementation of NVFP4 block-scaled GEMM.
fused-epilogueself.epi_tile = sm100_utils.compute_epilogue_tile_shape(
mbarriertmem_print_bar = pipeline.NamedBarrier(
num-warps = 1num_warps=1,
persistent-kernelProvides common interface for persistent tile scheduling.
shared-memorysmem_ptr: cute.Pointer, peer_cta_rank_in_cluster: Int32, *, loc=None, ip=None
tcgen05import cutlass.cute.nvgpu.tcgen05.mma as mma
tma"cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [$0], [$1], $2, [$3];",
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission.py3450 lines
import os
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import Optional, Tuple, Type, Union
from weakref import ref

import cuda.bindings.driver as cuda
import cutlass
import cutlass._mlir.dialects.cute as _cute_ir
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
import cutlass.cute as cute
import cutlass.cute.nvgpu.tcgen05 as tcgen05
import cutlass.cute.nvgpu.tcgen05.mma as mma
import cutlass.pipeline as pipeline
import cutlass.torch as cutlass_torch
import cutlass.utils as utils
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
import torch
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm
from cutlass.base_dsl import detect_gpu_arch
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.cute.runtime import from_dlpack, make_ptr
from cutlass.cutlass_dsl import (
    BFloat16,
    Boolean,
    Float4E2M1FN,
    Float8E4M3FN,
    Float8E5M2,
    Float16,
    Float32,
    Int8,
    Int32,
    Integer,
    Numeric,
    NumericMeta,
    T,
    TFloat32,
    Uint8,
    dsl_user_op,
    extract_mlir_values,
    min,
    new_from_mlir_values,
)
from cutlass.pipeline import (
    CooperativeGroup,
    PipelineAsync,
    PipelineOp,
    PipelineState,
    pipeline_init_arrive,
    pipeline_init_wait,
)
from task import input_t, output_t

DEBUG = os.environ.get("DEBUG_KERNEL", "0") == "1"

LINE_INFO = os.environ.get("LINE_INFO", "0") == "1"

N_TILE = int(os.environ.get("N_TILE", "64"))
assert N_TILE in [64, 128]

COMPILE_PTX = os.environ.get("COMPILE_PTX", "1") == "1"
DEBUG_BLOCK = (0, 0, 0)

USE_SIMPLE_SCHEDULER = os.environ.get("USE_SIMPLE_SCHEDULER", "1") == "1"

CLUSTER_N = int(os.environ.get("CLUSTER_N", "1"))
MAX_KSPLITS = int(os.environ.get("MAX_KSPLITS", "2"))
cluster_shape_mn = (1, CLUSTER_N)

# Kernel configuration parameters
# Tile sizes for M, N, K dimensions
mma_tiler_mnk = (128, N_TILE, 256)
mma_tiler_sfb_mnk = (128, 128, 256)
mma_inst_shape_mnk = (128, N_TILE, 64)
mma_inst_shape_mnk_sfb = (128, 128, 64)
mma_inst_tile_k = 4

# FP4 data type for A and B
ab_dtype = cutlass.Float4E2M1FN
# FP8 data type for scale factors
sf_dtype = cutlass.Float8E4M3FN
# FP16 output type
c_dtype = cutlass.Float16
# Scale factor block size (16 elements share one scale)
sf_vec_size = 16
# Number of threads per CUDA thread block
# Stage numbers of shared memory and tmem
num_acc_stage = 1
num_ab_stage = 1
# Total number of columns in tmem


@cute.jit
def is_debug_thr() -> bool:
    """Utility function to identify debug thread (thread 0 in block 0)

    :return: True if the current thread is the debug thread, False otherwise
    :rtype: bool
    """
    return cute.arch.thread_idx() == (0, 0, 0) and cute.arch.block_idx() == DEBUG_BLOCK


@cute.jit
def debug(*args):
    """Utility function to print debug information from the debug thread

    :param args: Arguments to print
    """
    if cutlass.const_expr(DEBUG):
        if is_debug_thr():
            cute.printf(*args)


@cute.jit
def hdebug(*args):
    """Utility function to print debug information from the debug thread

    :param args: Arguments to print
    """
    if cutlass.const_expr(DEBUG):
        cute.printf(*args)


@cute.jit
def alwaysprint(*args):
    """Utility function to always print debug information from all threads

    :param args: Arguments to print
    """
    if is_debug_thr():
        cute.printf(*args)


@cute.jit
def mma_debug(*args):
    """Utility function to print debug information from the debug thread

    :param args: Arguments to print
    """
    if cutlass.const_expr(DEBUG):
        if (
            cute.arch.block_idx() == DEBUG_BLOCK
            and (cute.arch.thread_idx()[0] % 32) == 0
        ):
            cute.printf(*args)


@cute.jit
def tma_debug(*args):
    mma_debug(*args)


@cute.jit
def print_tmem(
    tensor: cute.Tensor,
    num_cols: cutlass.Constexpr[cute.Int32],
    num_rows: cutlass.Constexpr[cute.Int32] = 32,
    num_warps=1,
):
    """Utility function to print tensor memory contents from the debug thread

    :param tensor: Tensor to print
    :type tensor: cute.Tensor
    """
    if cutlass.const_expr(True):
        tmem_print_bar = pipeline.NamedBarrier(
            barrier_id=10,
            num_threads=32 * num_warps,
        )
        tidx = cute.arch.thread_idx()[0]
        op = tcgen05.Ld32x32bOp(tcgen05.Repetition(num_cols), tcgen05.Pack.NONE)
        copy_atom_t2r = cute.make_copy_atom(op, cute.Float32)

        tmem_ptr = cute.recast_ptr(tensor.iterator, None, cute.Float32)
        tmem = cute.make_tensor(
            tmem_ptr, layout=cute.make_layout((32, num_cols), stride=(65536, 1))
        )

        tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tmem)
        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
        tTR_tTmem = thr_copy_t2r.partition_S(tmem)
        tTR_rTmem = cute.make_rmem_tensor((((1, num_cols), 1), 1, 1), cute.Float32)
        cute.copy(thr_copy_t2r, tTR_tTmem, tTR_rTmem)

        if tidx % 128 == 0:
            cute.printf("=== TMEM DUMP ===")
        for r in cutlass.range_constexpr(num_rows):
            if (tidx % 128) == r:
                cute.printf("row {}: {}", r, tTR_rTmem)
            tmem_print_bar.arrive_and_wait()
        if tidx % 128 == 0:
            cute.printf("=== TMEM DUMP END ===\n")


# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


def get_k_splits(k):
    ksplits = ceil_div(k, 4096)
    if ksplits > MAX_KSPLITS:
        ksplits = MAX_KSPLITS
    return ksplits


K_IDX = 0
N_IDX = 1


# borrowed from official flash attention repo
@dsl_user_op
def set_block_rank(
    smem_ptr: cute.Pointer, peer_cta_rank_in_cluster: Int32, *, loc=None, ip=None
) -> Int32:
    """Map the given smem pointer to the address at another CTA rank in the cluster."""
    smem_ptr_i32 = smem_ptr.toint(loc=loc, ip=ip).ir_value()
    return Int32(
        llvm.inline_asm(
            T.i32(),
            [smem_ptr_i32, peer_cta_rank_in_cluster.ir_value()],
            "mapa.shared::cluster.u32 $0, $1, $2;",
            "=r,r,r",
            has_side_effects=False,
            is_align_stack=False,
            asm_dialect=llvm.AsmDialect.AD_ATT,
        )
    )


@dsl_user_op
def store_shared_remote_fp32x4(
    a: Float32,
    b: Float32,
    c: Float32,
    d: Float32,
    smem_ptr: cute.Pointer,
    mbar_ptr: cute.Pointer,
    peer_cta_rank_in_cluster: Int32,
    *,
    loc=None,
    ip=None,
) -> None:
    remote_smem_ptr_i32 = set_block_rank(
        smem_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
    ).ir_value()
    remote_mbar_ptr_i32 = set_block_rank(
        mbar_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
    ).ir_value()
    llvm.inline_asm(
        None,
        [
            remote_smem_ptr_i32,
            remote_mbar_ptr_i32,
            Float32(a).ir_value(loc=loc, ip=ip),
            Float32(b).ir_value(loc=loc, ip=ip),
            Float32(c).ir_value(loc=loc, ip=ip),
            Float32(d).ir_value(loc=loc, ip=ip),
        ],
        "{\n\t"
        ".reg .v4 .f32 abcd;\n\t"
        "mov.f32 abcd.x, $2;\n\t"
        "mov.f32 abcd.y, $3;\n\t"
        "mov.f32 abcd.z, $4;\n\t"
        "mov.f32 abcd.w, $5;\n\t"
        "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"
        "}\n",
        "r,r,f,f,f,f",
        has_side_effects=True,
        is_align_stack=False,
        asm_dialect=llvm.AsmDialect.AD_ATT,
    )


@dsl_user_op
def cp_async_bulk_remote(
    dst_smem_ptr: cute.Pointer,
    src_smem_ptr: cute.Pointer,
    mbar_ptr: cute.Pointer,
    cp_bytes: cutlass.Int32,
    peer_cta_rank_in_cluster: Int32,
    *,
    loc=None,
    ip=None,
) -> None:
    src_smem_ptr_i32 = src_smem_ptr.toint(loc=loc, ip=ip).ir_value()
    remote_smem_ptr_i32 = set_block_rank(
        dst_smem_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
    ).ir_value()
    remote_mbar_ptr_i32 = set_block_rank(
        mbar_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
    ).ir_value()
    llvm.inline_asm(
        None,
        [
            remote_smem_ptr_i32,
            src_smem_ptr_i32,
            cute.Int32(cp_bytes).ir_value(loc=loc, ip=ip),
            remote_mbar_ptr_i32,
        ],
        "cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [$0], [$1], $2, [$3];",
        "r,r,r,r",
        has_side_effects=True,
        is_align_stack=False,
        asm_dialect=llvm.AsmDialect.AD_ATT,
    )


class WorkTileInfo:
    """A class to represent information about a work tile.

    :ivar tile_idx: The index of the tile.
    :type tile_idx: cute.Coord
    :ivar is_valid_tile: Whether the tile is valid.
    :type is_valid_tile: Boolean
    """

    def __init__(self, tile_idx: cute.Coord, is_valid_tile: Boolean):
        self._tile_idx = tile_idx
        self._is_valid_tile = Boolean(is_valid_tile)

    def __extract_mlir_values__(self) -> list[ir.Value]:
        values = extract_mlir_values(self.tile_idx)
        values.extend(extract_mlir_values(self.is_valid_tile))
        return values

    def __new_from_mlir_values__(self, values: list[ir.Value]) -> "WorkTileInfo":
        assert len(values) == 4
        new_tile_idx = new_from_mlir_values(self._tile_idx, values[:-1])
        new_is_valid_tile = new_from_mlir_values(self._is_valid_tile, [values[-1]])
        return WorkTileInfo(new_tile_idx, new_is_valid_tile)

    @property
    def is_valid_tile(self) -> Boolean:
        """Check latest tile returned by the scheduler is valid or not. Any scheduling
        requests after all tasks completed will return an invalid tile.

        :return: The validity of the tile.
        :rtype: Boolean
        """
        return self._is_valid_tile

    @property
    def tile_idx(self) -> cute.Coord:
        """
        Get the index of the tile.

        :return: The index of the tile.
        :rtype: cute.Coord
        """
        return self._tile_idx


@cute.jit
def get_mnkl_indices(
    tile: WorkTileInfo,
) -> tuple[
    cutlass.Constexpr[cute.Int32], cute.Int32, cute.Int32, cutlass.Constexpr[cute.Int32]
]:
    tile_idx = tile.tile_idx
    # tile_idx = (k_split, n_tile, 0)
    return 0, tile_idx[N_IDX], tile_idx[K_IDX], 0


class TileScheduler:
    """Base class for tile schedulers.

    Provides common interface for persistent tile scheduling.
    """

    @dsl_user_op
    def get_current_work(self, *, loc=None, ip=None) -> WorkTileInfo:
        raise NotImplementedError

    @dsl_user_op
    def initial_work_tile_info(self, *, loc=None, ip=None) -> WorkTileInfo:
        raise NotImplementedError

    @dsl_user_op
    def advance_to_next_work(self, *, advance_count: int = 1, loc=None, ip=None):
        raise NotImplementedError

    @property
    def num_tiles_executed(self) -> Int32:
        raise NotImplementedError

    @property
    def k_split_idx(self) -> Int32:
        """Returns the k-split index for this CTA."""
        raise NotImplementedError


class SimpleTileSchedulerParams:
    """Simplified tile scheduler parameters for fixed M=128, L=1 case.

    Grid layout: (k_splits, num_persistent_n_ctas, 1)
    - grid.x = k_splits (each x handles one K-split)
    - grid.y = min(N_tiles, max_clusters // k_splits)

    Work tile: (k_split, n_tile) where:
    - k_split = bidx (fixed for CTA lifetime)
    - n_tile = bidy + wave * n_stride (advances each iteration)
    """

    def __init__(
        self,
        k_splits: int,
        n_tiles: int,
        n_stride: int,
        *,
        loc=None,
        ip=None,
    ):
        self.k_splits = k_splits
        self.n_tiles = n_tiles
        self.n_stride = n_stride
        self._loc = loc

    def __extract_mlir_values__(self):
        # k_splits, n_tiles, and n_stride are static Python ints, no MLIR values to extract
        return []

    def __new_from_mlir_values__(self, values):
        # No MLIR values to reconstruct, just return a copy with same static values
        return SimpleTileSchedulerParams(
            self.k_splits,
            self.n_tiles,
            self.n_stride,
            loc=self._loc,
        )

    @staticmethod
    def get_grid_shape(
        cluster_shape_kn: Tuple[int, int],
        n_tiles: int,
        max_active_clusters: int,
        *,
        loc=None,
        ip=None,
    ) -> Tuple[int, int, int]:
        """Computes grid shape: (k_splits, num_persistent_n_ctas, 1)"""
        # Number of N-tile CTAs we can run per wave
        ksplits, cluster_n = cluster_shape_kn
        max_n_ctas = max_active_clusters * cluster_n
        num_persistent_n_ctas = min(n_tiles, max_n_ctas)
        return (ksplits, num_persistent_n_ctas, 1)


class SimpleTileScheduler(TileScheduler):
    """Simplified persistent tile scheduler for fixed M=128, L=1 case.

    Each CTA has:
    - Fixed k_split = block_idx.x
    - Advancing n_tile = block_idx.y + wave * grid_dim.y
    """

    def __init__(
        self,
        params: SimpleTileSchedulerParams,
        k_split: Int32,
        current_n_tile: Int32,
        num_tiles_executed: Int32,
    ):
        self.params = params
        self.k_split = k_split
        self._current_n_tile = current_n_tile
        self._num_tiles_executed = num_tiles_executed

    def __extract_mlir_values__(self) -> list[ir.Value]:
        values = extract_mlir_values(self.k_split)
        values.extend(extract_mlir_values(self._current_n_tile))
        values.extend(extract_mlir_values(self._num_tiles_executed))
        return values

    def __new_from_mlir_values__(self, values: list[ir.Value]) -> "SimpleTileScheduler":
        assert len(values) == 3
        return SimpleTileScheduler(
            self.params,
            new_from_mlir_values(self.k_split, [values[0]]),
            new_from_mlir_values(self._current_n_tile, [values[1]]),
            new_from_mlir_values(self._num_tiles_executed, [values[2]]),
        )

    @staticmethod
    @dsl_user_op
    def create(
        params: SimpleTileSchedulerParams,
        block_idx: Tuple[Integer, Integer, Integer],
        *,
        loc=None,
        ip=None,
    ):
        """Initialize the simplified tile scheduler.

        block_idx.x = k_split (fixed)
        block_idx.y = initial n_tile offset
        params.n_stride = how much to advance each wave (statically known)
        """
        bidx, bidy, _ = block_idx

        k_split = Int32(bidx)
        current_n_tile = Int32(bidy)
        num_tiles_executed = Int32(0)

        return SimpleTileScheduler(params, k_split, current_n_tile, num_tiles_executed)

    @dsl_user_op
    def get_current_work(self, *, loc=None, ip=None) -> WorkTileInfo:
        """Returns current work tile: (k_split, n_tile, 0)"""
        is_valid = self._current_n_tile < Int32(self.params.n_tiles)
        tile_idx = (self.k_split, self._current_n_tile, Int32(0))
        return WorkTileInfo(tile_idx, is_valid)

    @dsl_user_op
    def initial_work_tile_info(self, *, loc=None, ip=None) -> WorkTileInfo:
        return self.get_current_work(loc=loc, ip=ip)

    @dsl_user_op
    def advance_to_next_work(self, *, advance_count: int = 1, loc=None, ip=None):
        """Advance to next N-tile: n_tile += n_stride"""
        self._current_n_tile += Int32(advance_count * self.params.n_stride)
        self._num_tiles_executed += Int32(1)

    @property
    def num_tiles_executed(self) -> Int32:
        return self._num_tiles_executed

    @property
    def k_split_idx(self) -> Int32:
        """Returns the k-split index for this CTA (block_idx.x)."""
        return self.k_split


class PersistentTileSchedulerParams:
    """A class to represent parameters for a persistent tile scheduler.

    This class is designed to manage and compute the layout of clusters and tiles
    in a batched gemm problem.

    :ivar cluster_shape_mn: Shape of the cluster in (m, n) dimensions (K dimension cta count must be 1).
    :type cluster_shape_mn: tuple
    :ivar problem_layout_ncluster_mnl: Layout of the problem in terms of
        number of clusters in (m, n, l) dimensions.
    :type problem_layout_ncluster_mnl: cute.Layout
    """

    def __init__(
        self,
        problem_shape_ntile_mnl: cute.Shape,
        cluster_shape_mnk: cute.Shape,
        swizzle_size: int = 1,
        raster_along_m: bool = True,
        *,
        loc=None,
        ip=None,
    ):
        """
        Initializes the PersistentTileSchedulerParams with the given parameters.

        :param problem_shape_ntile_mnl: The shape of the problem in terms of
            number of CTA (Cooperative Thread Array) in (m, n, l) dimensions.
        :type problem_shape_ntile_mnl: cute.Shape
        :param cluster_shape_mnk: The shape of the cluster in (m, n) dimensions.
        :type cluster_shape_mnk: cute.Shape
        :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle
        :type swizzle_size: int
        :param raster_along_m: Rasterization order of clusters. Only used when swizzle_size > 1.
            True means along M, false means along N.
        :type raster_along_m: bool

        :raises ValueError: If cluster_shape_k is not 1.
        """

        # if cluster_shape_mnk[2] != 1:
        #     raise ValueError(f"unsupported cluster_shape_k {cluster_shape_mnk[2]}")
        if swizzle_size < 1:
            raise ValueError(f"expect swizzle_size >= 1, but get {swizzle_size}")

        self.problem_shape_ntile_mnl = problem_shape_ntile_mnl
        # cluster_shape_mnk is kept for reconstruction
        self._cluster_shape_mnk = cluster_shape_mnk
        self.cluster_shape_mn = cluster_shape_mnk[:2]
        self.swizzle_size = swizzle_size
        self._raster_along_m = raster_along_m
        self._loc = loc

        # By default, we follow m major (col-major) raster order, so make a col-major layout
        self.problem_layout_ncluster_mnl = cute.make_layout(
            cute.ceil_div(
                self.problem_shape_ntile_mnl, cluster_shape_mnk[:2], loc=loc, ip=ip
            ),
            loc=loc,
            ip=ip,
        )

        if swizzle_size > 1:
            problem_shape_ncluster_mnl = cute.round_up(
                self.problem_layout_ncluster_mnl.shape,
                (1, swizzle_size, 1) if raster_along_m else (swizzle_size, 1, 1),
            )

            if raster_along_m:
                self.problem_layout_ncluster_mnl = cute.make_layout(
                    (
                        problem_shape_ncluster_mnl[0],
                        (swizzle_size, problem_shape_ncluster_mnl[1] // swizzle_size),
                        problem_shape_ncluster_mnl[2],
                    ),
                    stride=(
                        swizzle_size,
                        (1, swizzle_size * problem_shape_ncluster_mnl[0]),
                        problem_shape_ncluster_mnl[0] * problem_shape_ncluster_mnl[1],
                    ),
                    loc=loc,
                    ip=ip,
                )
            else:
                self.problem_layout_ncluster_mnl = cute.make_layout(
                    (
                        (swizzle_size, problem_shape_ncluster_mnl[0] // swizzle_size),
                        problem_shape_ncluster_mnl[1],
                        problem_shape_ncluster_mnl[2],
                    ),
                    stride=(
                        (1, swizzle_size * problem_shape_ncluster_mnl[1]),
                        swizzle_size,
                        problem_shape_ncluster_mnl[0] * problem_shape_ncluster_mnl[1],
                    ),
                    loc=loc,
                    ip=ip,
                )

    def __extract_mlir_values__(self):
        values, self._values_pos = [], []
        for obj in [
            self.problem_shape_ntile_mnl,
            self._cluster_shape_mnk,
            self.swizzle_size,
            self._raster_along_m,
        ]:
            obj_values = extract_mlir_values(obj)
            values += obj_values
            self._values_pos.append(len(obj_values))
        return values

    def __new_from_mlir_values__(self, values):
        obj_list = []
        for obj, n_items in zip(
            [
                self.problem_shape_ntile_mnl,
                self._cluster_shape_mnk,
                self.swizzle_size,
                self._raster_along_m,
            ],
            self._values_pos,
        ):
            obj_list.append(new_from_mlir_values(obj, values[:n_items]))
            values = values[n_items:]
        return PersistentTileSchedulerParams(*(tuple(obj_list)), loc=self._loc)

    @dsl_user_op
    def get_grid_shape(
        self, max_active_clusters: Int32, *, loc=None, ip=None
    ) -> Tuple[Integer, Integer, Integer]:
        """
        Computes the grid shape based on the maximum active clusters allowed.

        :param max_active_clusters: The maximum number of active clusters that
            can run in one wave.
        :type max_active_clusters: Int32

        :return: A tuple containing the grid shape in (m, n, persistent_clusters).
            - m: self.cluster_shape_m.
            - n: self.cluster_shape_n.
            - persistent_clusters: Number of persistent clusters that can run.
        """

        # Total ctas in problem size
        num_ctas_mnl = tuple(
            cute.size(x) * y
            for x, y in zip(
                self.problem_layout_ncluster_mnl.shape, self.cluster_shape_mn
            )
        ) + (self.problem_layout_ncluster_mnl.shape[2],)

        num_ctas_in_problem = cute.size(num_ctas_mnl, loc=loc, ip=ip)

        num_ctas_per_cluster = cute.size(self.cluster_shape_mn, loc=loc, ip=ip)
        # Total ctas that can run in one wave
        num_ctas_per_wave = max_active_clusters * num_ctas_per_cluster

        num_persistent_ctas = min(num_ctas_in_problem, num_ctas_per_wave)
        num_persistent_clusters = num_persistent_ctas // num_ctas_per_cluster

        return (*self.cluster_shape_mn, num_persistent_clusters)


class StaticPersistentTileScheduler(TileScheduler):
    """A scheduler for static persistent tile execution in CUTLASS/CuTe kernels.

    :ivar params: Tile schedule related params, including cluster shape and problem_layout_ncluster_mnl
    :type params: PersistentTileSchedulerParams
    :ivar num_persistent_clusters: Number of persistent clusters that can be launched
    :type num_persistent_clusters: Int32
    :ivar cta_id_in_cluster: ID of the CTA within its cluster
    :type cta_id_in_cluster: cute.Coord
    :ivar _num_tiles_executed: Counter for executed tiles
    :type _num_tiles_executed: Int32
    :ivar _current_work_linear_idx: Current cluster index
    :type _current_work_linear_idx: Int32
    """

    def __init__(
        self,
        params: PersistentTileSchedulerParams,
        num_persistent_clusters: Int32,
        current_work_linear_idx: Int32,
        cta_id_in_cluster: cute.Coord,
        num_tiles_executed: Int32,
    ):
        """
        Initializes the StaticPersistentTileScheduler with the given parameters.

        :param params: Tile schedule related params, including cluster shape and problem_layout_ncluster_mnl.
        :type params: PersistentTileSchedulerParams
        :param num_persistent_clusters: Number of persistent clusters that can be launched.
        :type num_persistent_clusters: Int32
        :param current_work_linear_idx: Current cluster index.
        :type current_work_linear_idx: Int32
        :param cta_id_in_cluster: ID of the CTA within its cluster.
        :type cta_id_in_cluster: cute.Coord
        :param num_tiles_executed: Counter for executed tiles.
        :type num_tiles_executed: Int32
        """
        self.params = params
        self.num_persistent_clusters = num_persistent_clusters
        self._current_work_linear_idx = current_work_linear_idx
        self.cta_id_in_cluster = cta_id_in_cluster
        self._num_tiles_executed = num_tiles_executed

    def __extract_mlir_values__(self) -> list[ir.Value]:
        values = extract_mlir_values(self.num_persistent_clusters)
        values.extend(extract_mlir_values(self._current_work_linear_idx))
        values.extend(extract_mlir_values(self.cta_id_in_cluster))
        values.extend(extract_mlir_values(self._num_tiles_executed))
        return values

    def __new_from_mlir_values__(
        self, values: list[ir.Value]
    ) -> "StaticPersistentTileScheduler":
        assert len(values) == 6
        new_num_persistent_clusters = new_from_mlir_values(
            self.num_persistent_clusters, [values[0]]
        )
        new_current_work_linear_idx = new_from_mlir_values(
            self._current_work_linear_idx, [values[1]]
        )
        new_cta_id_in_cluster = new_from_mlir_values(
            self.cta_id_in_cluster, values[2:5]
        )
        new_num_tiles_executed = new_from_mlir_values(
            self._num_tiles_executed, [values[5]]
        )
        return StaticPersistentTileScheduler(
            self.params,
            new_num_persistent_clusters,
            new_current_work_linear_idx,
            new_cta_id_in_cluster,
            new_num_tiles_executed,
        )

    @staticmethod
    @dsl_user_op
    def create(
        params: PersistentTileSchedulerParams,
        block_idx: Tuple[Integer, Integer, Integer],
        grid_dim: Tuple[Integer, Integer, Integer],
        *,
        loc=None,
        ip=None,
    ):
        """Initialize the static persistent tile scheduler.

        :param params: Parameters for the persistent
            tile scheduler.
        :type params: PersistentTileSchedulerParams
        :param block_idx: The 3d block index in the format (bidx, bidy, bidz).
        :type block_idx: Tuple[Integer, Integer, Integer]
        :param grid_dim: The 3d grid dimensions for kernel launch.
        :type grid_dim: Tuple[Integer, Integer, Integer]

        :return: A StaticPersistentTileScheduler object.
        :rtype: StaticPersistentTileScheduler
        """

        # Calculate the number of persistent clusters by dividing the total grid size
        # by the number of CTAs per cluster
        num_persistent_clusters = cute.size(grid_dim, loc=loc, ip=ip) // cute.size(
            params.cluster_shape_mn, loc=loc, ip=ip
        )

        bidx, bidy, bidz = block_idx

        # Initialize workload index equals to the cluster index in the grid
        current_work_linear_idx = Int32(bidz)

        # CTA id in the cluster
        cta_id_in_cluster = (
            Int32(bidx % params.cluster_shape_mn[0]),
            Int32(bidy % params.cluster_shape_mn[1]),
            Int32(0),
        )
        # Initialize number of tiles executed to zero
        num_tiles_executed = Int32(0)
        return StaticPersistentTileScheduler(
            params,
            num_persistent_clusters,
            current_work_linear_idx,
            cta_id_in_cluster,
            num_tiles_executed,
        )

    # called by host
    @staticmethod
    def get_grid_shape(
        params: PersistentTileSchedulerParams,
        max_active_clusters: Int32,
        *,
        loc=None,
        ip=None,
    ) -> Tuple[Integer, Integer, Integer]:
        """Calculates the grid shape to be launched on GPU using problem shape,
        threadblock shape, and active cluster size.

        :param params: Parameters for grid shape calculation.
        :type params: PersistentTileSchedulerParams
        :param max_active_clusters: Maximum active clusters allowed.
        :type max_active_clusters: Int32

        :return: The calculated 3d grid shape.
        :rtype: Tuple[Integer, Integer, Integer]
        """

        return params.get_grid_shape(max_active_clusters, loc=loc, ip=ip)

    # private method
    def _get_current_work_for_linear_idx(
        self, current_work_linear_idx: Int32, *, loc=None, ip=None
    ) -> WorkTileInfo:
        """Compute current tile coord given current_work_linear_idx and cta_id_in_cluster.

        :param current_work_linear_idx: The linear index of the current work.
        :type current_work_linear_idx: Int32

        :return: An object containing information about the current tile coordinates
            and validity status.
        :rtype: WorkTileInfo
        """

        is_valid = current_work_linear_idx < cute.size(
            self.params.problem_layout_ncluster_mnl, loc=loc, ip=ip
        )

        if self.params.swizzle_size == 1:
            cur_cluster_coord = self.params.problem_layout_ncluster_mnl.get_hier_coord(
                current_work_linear_idx, loc=loc, ip=ip
            )
        else:
            cur_cluster_coord = self.params.problem_layout_ncluster_mnl.get_flat_coord(
                current_work_linear_idx, loc=loc, ip=ip
            )

        # cur_tile_coord is a tuple of i32 values
        cur_tile_coord = tuple(
            Int32(x) * Int32(z) + Int32(y)
            for x, y, z in zip(
                cur_cluster_coord,
                self.cta_id_in_cluster,
                (*self.params.cluster_shape_mn, Int32(1)),
            )
        )

        return WorkTileInfo(cur_tile_coord, is_valid)

    @dsl_user_op
    def get_current_work(self, *, loc=None, ip=None) -> WorkTileInfo:
        return self._get_current_work_for_linear_idx(
            self._current_work_linear_idx, loc=loc, ip=ip
        )

    @dsl_user_op
    def initial_work_tile_info(self, *, loc=None, ip=None) -> WorkTileInfo:
        return self.get_current_work(loc=loc, ip=ip)

    @dsl_user_op
    def advance_to_next_work(self, *, advance_count: int = 1, loc=None, ip=None):
        self._current_work_linear_idx += Int32(advance_count) * Int32(
            self.num_persistent_clusters
        )
        self._num_tiles_executed += Int32(1)

    @property
    def num_tiles_executed(self) -> Int32:
        return self._num_tiles_executed

    @property
    def k_split_idx(self) -> Int32:
        """Returns the k-split index for this CTA (cta_id_in_cluster[0] when using cluster_shape_kn)."""
        return self.cta_id_in_cluster[0]


# Type alias for scheduler params - either SimpleTileSchedulerParams or PersistentTileSchedulerParams
TileSchedulerParams = Union[SimpleTileSchedulerParams, PersistentTileSchedulerParams]


def create_tile_scheduler(
    tile_sched_params: TileSchedulerParams,
    block_idx: Tuple[Integer, Integer, Integer],
    grid_dim: Tuple[Integer, Integer, Integer],
) -> TileScheduler:
    """Factory function to create the appropriate tile scheduler based on USE_SIMPLE_SCHEDULER flag."""
    if cutlass.const_expr(USE_SIMPLE_SCHEDULER):
        return SimpleTileScheduler.create(tile_sched_params, block_idx)
    else:
        return StaticPersistentTileScheduler.create(
            tile_sched_params, block_idx, grid_dim
        )


@cute.jit
def get_nk_coords(tile_idx: cute.Coord) -> Tuple[cute.Int32, cute.Int32]:
    pass


class FP4GEMMKernel:
    def __init__(
        self,
    ):
        self.acc_dtype = cutlass.Float32
        self.ab_dtype = ab_dtype
        self.sf_vec_size = sf_vec_size
        self.sf_dtype = sf_dtype
        self.c_dtype = c_dtype
        # only used for reductions.
        # TODO: refactor, it's currently confusing
        self.d_dtype = c_dtype
        self.cluster_shape_mn = cluster_shape_mn
        self.mma_tiler = mma_tiler_mnk
        self.mma_tiler_sfb = mma_tiler_sfb_mnk

        self.major_mode = tcgen05.OperandMajorMode.K
        self.cd_majorness = utils.LayoutEnum.ROW_MAJOR
        self.cta_group = tcgen05.CtaGroup.ONE

        self.threads_per_warp = 32
        self.occupancy = 1
        # Set specialized warp ids
        self.epilog_warp_id = (
            0,
            1,
            2,
            3,
        )
        self.mma_warp_id = 4
        self.tma_warp_id = 5
        self.threads_per_cta = self.threads_per_warp * len(
            (self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
        )
        # Set barrier id for epilogue sync and tmem ptr sync
        self.epilog_sync_barrier = pipeline.NamedBarrier(
            barrier_id=1,
            num_threads=self.threads_per_warp * len(self.epilog_warp_id),
        )
        self.tmem_alloc_barrier = pipeline.NamedBarrier(
            barrier_id=2,
            num_threads=self.threads_per_warp
            * len((self.mma_warp_id, *self.epilog_warp_id)),
        )
        self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
        SM100_TMEM_CAPACITY_COLUMNS = 512
        self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS

    @cute.jit
    def is_work_tile_valid(self, work_tile: WorkTileInfo) -> Boolean:
        """Check if the work tile is valid.

        :param work_tile: The work tile to check.
        :return: True if the work tile is valid, False otherwise.
        """
        # return work_tile.tile_idx[N_IDX] == cute.arch.block_idx()[1]
        return work_tile.is_valid_tile

    @cute.jit
    def __call__(
        self,
        a_tensor: cute.Tensor,
        b_tensor: cute.Tensor,
        sfa_tensor: cute.Tensor,
        sfb_tensor: cute.Tensor,
        c_tensor: cute.Tensor,
        max_active_clusters: cutlass.Constexpr,
    ):
        hdebug("\n===================== Input Tensors =====================")
        hdebug("a_tensor layout: {}", a_tensor.layout)
        hdebug("b_tensor layout: {}", b_tensor.layout)
        hdebug("sfa_tensor layout: {}", sfa_tensor.layout)
        hdebug("sfb_tensor layout: {}", sfb_tensor.layout)
        hdebug("c_tensor layout: {}", c_tensor.layout)
        hdebug("==========================================================\n\n")

        tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.ab_dtype,
            self.major_mode,
            self.major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            self.cta_group,
            self.mma_tiler[:2],
        )

        tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
            self.ab_dtype,
            self.major_mode,
            self.major_mode,
            self.sf_dtype,
            self.sf_vec_size,
            self.cta_group,
            self.mma_tiler_sfb[:2],
        )

        self.ksplits, self.ksplits_tile_cnt = self._get_k_splits(a_tensor)
        assert cute.is_static(self.ksplits_tile_cnt)
        self.cluster_shape_kn = (self.ksplits, self.cluster_shape_mn[1])

        # Compute 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,),
        )
        self.cluster_layout_red_vmnk = cute.tiled_divide(
            cute.make_layout((*self.cluster_shape_kn, 1)),
            (tiled_mma_sfb.thr_id.shape,),
        )

        # Compute number of multicast CTAs for A/B
        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.is_a_mcast = self.num_mcast_ctas_a > 1

        # N_tiles should be based on mma_tiler[1], not SFB tiles (which are always 128)
        # When N_TILE=64, there are 2x as many MMA tiles as SFB tiles
        # b_tensor shape is (N, K, L), so N is at mode 0
        N = cute.size(b_tensor, mode=[0])
        N_tiles = N // self.mma_tiler[1]

        # Compute epilogue subtile
        self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
            self.mma_tiler,
            False,
            self.cd_majorness,
            self.c_dtype,
        )
        if cutlass.const_expr(self.ksplits > 1):
            self.epi_tile: cute.Tile = (
                self.mma_tiler[0],
                self.mma_tiler[1] // self.ksplits,
            )
            self.c_dtype = cute.Float32
        self.epi_tile_n = cute.size(self.epi_tile[1])

        # Setup A/B/C stage count in shared memory and ACC stage count in tensor memory
        self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
            tiled_mma,
            self.mma_tiler,
            self.ab_dtype,
            self.ab_dtype,
            self.epi_tile,
            self.c_dtype,
            self.cd_majorness,
            self.sf_dtype,
            self.sf_vec_size,
            self.smem_capacity,
        )

        # Compute A/B/SFA/SFB/C shared memory layout
        self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(
            tiled_mma,
            self.mma_tiler,
            self.ab_dtype,
            self.num_ab_stage,
        )
        self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(
            tiled_mma,
            self.mma_tiler,
            self.ab_dtype,
            self.num_ab_stage,
        )
        self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
        )
        self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            self.mma_tiler,
            self.sf_vec_size,
            self.num_ab_stage,
        )
        self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
            self.c_dtype,
            self.cd_majorness,
            self.epi_tile,
            self.num_c_stage,
        )

        hdebug("\n===================== SMEM Layouts (staged) =====================")
        hdebug("a_smem_layout_staged: {}", self.a_smem_layout_staged)
        hdebug("b_smem_layout_staged: {}", self.b_smem_layout_staged)
        hdebug("sfa_smem_layout_staged: {}", self.sfa_smem_layout_staged)
        hdebug("sfb_smem_layout_staged: {}", self.sfb_smem_layout_staged)
        hdebug("c_smem_layout_staged: {}", self.c_smem_layout_staged)
        hdebug("==================================================================\n\n")

        # Compute number of TMEM columns for SFA/SFB/Accumulator
        sf_atom_mn = 32
        self.num_sfa_tmem_cols = (self.mma_tiler[0] // sf_atom_mn) * mma_inst_tile_k
        self.num_sfb_tmem_cols = (self.mma_tiler_sfb[1] // sf_atom_mn) * mma_inst_tile_k
        self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
        self.num_accumulator_tmem_cols = self.mma_tiler[1] * self.num_acc_stage

        # Setup sfa/sfb tensor by filling A/B tensor to scale factor atom layout
        # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL)
        sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
            a_tensor.shape, self.sf_vec_size
        )
        sfa_tensor = cute.make_tensor(sfa_tensor.iterator, sfa_layout)

        # ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL)
        sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
            b_tensor.shape, self.sf_vec_size
        )
        sfb_tensor = cute.make_tensor(sfb_tensor.iterator, sfb_layout)
        atom_thr_size = cute.size(tiled_mma.thr_id.shape)

        # Setup TMA load for A
        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,
        )

        # Setup TMA load for B
        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,
        )

        # Setup TMA load for SFA
        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,
        )

        # Setup TMA load for SFB
        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.ab_dtype, a_smem_layout)
        b_copy_size = cute.size_in_bytes(self.ab_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

        hdebug("\n===================== TMA Atoms and Tensors =====================")
        hdebug(f"\ntma_atom_a:\n{tma_atom_a}\n")
        hdebug("tma_tensor_a layout: {}", tma_tensor_a.layout)
        hdebug("a_copy_size: {}", a_copy_size)
        hdebug(f"\ntma_atom_b:\n{tma_atom_b}\n")
        hdebug("tma_tensor_b layout: {}", tma_tensor_b.layout)
        hdebug("b_copy_size: {}", b_copy_size)
        hdebug(f"\ntma_atom_sfa:\n{tma_atom_sfa}\n")
        hdebug("tma_tensor_sfa layout: {}", tma_tensor_sfa.layout)
        hdebug("sfa_copy_size: {}", sfa_copy_size)
        hdebug(f"\ntma_atom_sfb:\n{tma_atom_sfb}\n")
        hdebug("tma_tensor_sfb layout: {}", tma_tensor_sfb.layout)
        hdebug("sfb_copy_size: {}", sfb_copy_size)
        hdebug("num_tma_load_bytes: {}", self.num_tma_load_bytes)
        hdebug("==================================================================\n\n")

        # Setup TMA store for C
        epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
        c_copy_op = (
            cpasync.CopyBulkTensorTileS2GOp()
            if self.ksplits == 1
            else cpasync.CopyReduceBulkTensorTileS2GOp(
                reduction_kind=cpasync.ReductionOp.ADD
            )
        )
        tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
            c_copy_op,
            c_tensor,
            epi_smem_layout,
            self.epi_tile,
        )

        hdebug("\n===================== TMA Store (C) =====================")
        hdebug(f"\ntma_atom_c:\n{tma_atom_c}\n")
        hdebug("tma_tensor_c layout: {}", tma_tensor_c.layout)
        hdebug("N_tiles: {}", N_tiles)
        hdebug("==========================================================\n\n")

        # Compute grid size
        self.tile_sched_params, grid = self._compute_grid(
            N_tiles,
            self.cluster_shape_kn,
            max_active_clusters,
        )

        self.buffer_align_bytes = 1024

        num_red_stages = 0
        self.red_recv_bytes = cute.Int32(0)
        self.red_send_bytes = cute.Int32(0)
        if cutlass.const_expr(self.ksplits > 1):
            num_red_stages = self.num_acc_stage
            self.red_send_bytes = cute.size_in_bytes(
                self.c_dtype,
                cute.make_layout(self.epi_tile),
            )
            self.red_send_bytes_per_warp = self.red_send_bytes // len(
                self.epilog_warp_id
            )
            self.red_recv_bytes = self.red_send_bytes * (self.ksplits - 1)

        hdebug("\n===================== Kernel Shared Storage =====================")
        hdebug("num_ab_stage: {}", self.num_ab_stage)
        hdebug("num_acc_stage: {}", self.num_acc_stage)
        hdebug("num_red_stages: {}", num_red_stages)
        hdebug("red_recv_bytes: {}", self.red_recv_bytes)
        hdebug("==================================================================\n\n")

        # Define shared storage for kernel
        @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]
            red_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_red_stages]
            red_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_red_stages]
            tmem_dealloc_mbar_ptr: cutlass.Int64
            tmem_holding_buf: cutlass.Int32
            # (EPI_TILE_M, EPI_TILE_N, STAGE)
            sC: cute.struct.Align[
                cute.struct.MemRange[
                    self.c_dtype,
                    cute.cosize(self.c_smem_layout_staged.outer),
                ],
                self.buffer_align_bytes,
            ]
            # (MMA, MMA_M, MMA_K, STAGE)
            sA: cute.struct.Align[
                cute.struct.MemRange[
                    self.ab_dtype, cute.cosize(self.a_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]
            # (MMA, MMA_N, MMA_K, STAGE)
            sB: cute.struct.Align[
                cute.struct.MemRange[
                    self.ab_dtype, cute.cosize(self.b_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]
            # (MMA, MMA_M, MMA_K, STAGE)
            sSFA: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]
            # (MMA, MMA_N, MMA_K, STAGE)
            sSFB: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]

        self.shared_storage = SharedStorage

        # Launch the kernel synchronously
        self.kernel(
            tiled_mma,
            tiled_mma_sfb,
            tma_atom_a,
            tma_tensor_a,
            tma_atom_b,
            tma_tensor_b,
            tma_atom_sfa,
            tma_tensor_sfa,
            tma_atom_sfb,
            tma_tensor_sfb,
            tma_atom_c,
            tma_tensor_c,
            c_tensor,
            self.cluster_layout_vmnk,
            self.cluster_layout_sfb_vmnk,
            self.cluster_layout_red_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,
            self.tile_sched_params,
        ).launch(
            grid=grid,
            block=[self.threads_per_cta, 1, 1],
            cluster=(*self.cluster_shape_kn, 1),
            min_blocks_per_mp=1,
        )

    # GPU device kernel
    @cute.kernel
    def kernel(
        self,
        tiled_mma: cute.TiledMma,
        tiled_mma_sfb: cute.TiledMma,
        tma_atom_a: cute.CopyAtom,
        mA_mkl: cute.Tensor,
        tma_atom_b: cute.CopyAtom,
        mB_nkl: cute.Tensor,
        tma_atom_sfa: cute.CopyAtom,
        mSFA_mkl: cute.Tensor,
        tma_atom_sfb: cute.CopyAtom,
        mSFB_nkl: cute.Tensor,
        tma_atom_c: cute.CopyAtom,
        mC_mnl: cute.Tensor,
        mC_raw: cute.Tensor,
        cluster_layout_vmnk: cute.Layout,
        cluster_layout_sfb_vmnk: cute.Layout,
        cluster_layout_red_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: TileSchedulerParams,
    ):
        """
        GPU device kernel performing the Persistent batched GEMM computation.
        """
        warp_idx = cute.arch.warp_idx()
        warp_idx = cute.arch.make_warp_uniform(warp_idx)

        #
        # Prefetch tma desc
        #
        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)
        # if cutlass.const_expr(self.ksplits == 1):
        #     if warp_idx == self.epilog_warp_id[3]:
        #         cpasync.prefetch_descriptor(tma_atom_c)

        #
        # Setup cta/thread coordinates
        #
        # Coords inside cluster
        cta_rank_in_cluster = cute.arch.make_warp_uniform(
            cute.arch.block_idx_in_cluster()
        )
        block_in_cluster_coord_vmnk = cluster_layout_red_vmnk.get_flat_coord(
            cta_rank_in_cluster
        )
        # Coord inside cta
        tidx, _, _ = cute.arch.thread_idx()

        #
        # Alloc and init: a+b full/empty, accumulator full/empty, tensor memory dealloc barrier
        #
        smem = utils.SmemAllocator()
        storage = smem.allocate(self.shared_storage)

        # Initialize mainloop ab_pipeline (barrier) and states
        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_red_vmnk,
            mcast_mode_mn=(1, 0),
            defer_sync=True,
        )

        # Initialize acc_pipeline (barrier) and states
        acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_acc_consumer_threads = len(self.epilog_warp_id)
        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,
        )

        red_pipeline = None
        if cutlass.const_expr(self.ksplits > 1):
            red_pipeline_producer_group = pipeline.CooperativeGroup(
                pipeline.Agent.Thread
            )
            red_pipeline_consumer_group = pipeline.CooperativeGroup(
                pipeline.Agent.Thread, self.ksplits * len(self.epilog_warp_id)
            )
            red_pipeline = pipeline.PipelineTmaAsync.create(
                num_stages=self.num_acc_stage,
                producer_group=red_pipeline_producer_group,
                consumer_group=red_pipeline_consumer_group,
                tx_count=self.red_recv_bytes,
                barrier_storage=storage.red_full_mbar_ptr.data_ptr(),
                cta_layout_vmnk=cluster_layout_red_vmnk,
                tidx=tidx,
                # Multicast along N dimension (mode 2) so that CTAs with the same
                # cluster N index communicate: ranks 0↔1 and ranks 2↔3 exchange data
                # cluster_layout_red_vmnk has shape (v, ksplits, N, k) = (1, 2, 2, 1)
                mcast_mode_mn=(0, 1),
                defer_sync=True,
            )

        # Tensor memory dealloc barrier init
        tmem = utils.TmemAllocator(
            storage.tmem_holding_buf,
            barrier_for_retrieve=self.tmem_alloc_barrier,
            allocator_warp_id=self.epilog_warp_id[0],
        )

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

        #
        # Compute multicast mask for A/B/SFA/SFB buffer full
        #
        a_full_mcast_mask = None
        sfa_full_mcast_mask = None
        sfb_full_mcast_mask = None
        if cutlass.const_expr(self.is_a_mcast):
            a_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_red_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )
            sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_red_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )

        #
        # Setup smem tensor A/B/SFA/SFB/C
        #
        # (EPI_TILE_M, EPI_TILE_N, STAGE)
        sC = storage.sC.get_tensor(
            c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
        )
        # (MMA, MMA_M, MMA_K, STAGE)
        sA = storage.sA.get_tensor(
            a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
        )
        # (MMA, MMA_N, MMA_K, STAGE)
        sB = storage.sB.get_tensor(
            b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
        )
        # (MMA, MMA_M, MMA_K, STAGE)
        sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
        # (MMA, MMA_N, MMA_K, STAGE)
        sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)

        debug("\n===================== SMEM Tensors =====================")
        debug("sC layout: {}", sC.layout)
        debug("sA layout: {}", sA.layout)
        debug("sB layout: {}", sB.layout)
        debug("sSFA layout: {}", sSFA.layout)
        debug("sSFB layout: {}", sSFB.layout)
        debug("========================================================\n\n")

        #
        # Local_tile partition global tensors
        #
        # (bM, bK, RestM, RestK, RestL)
        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        # (bN, bK, RestN, RestK, RestL)
        gB_nkl = cute.local_tile(
            mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
        )
        # (bM, bK, RestM, RestK, RestL)
        gSFA_mkl = cute.local_tile(
            mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )
        # (bN, bK, RestN, RestK, RestL)
        gSFB_nkl = cute.local_tile(
            mSFB_nkl,
            cute.slice_(self.mma_tiler_sfb, (0, None, None)),
            (None, None, None),
        )
        # (bM, bN, RestM, RestN, RestL)
        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
        )

        gC_raw = cute.local_tile(
            mC_raw, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
        )

        debug(
            "\n===================== Global Tensors (local_tile) ====================="
        )
        debug("gA_mkl layout: {}", gA_mkl.layout)
        debug("gB_nkl layout: {}", gB_nkl.layout)
        debug("gSFA_mkl layout: {}", gSFA_mkl.layout)
        debug("gSFB_nkl layout: {}", gSFB_nkl.layout)
        debug("gC_mnl layout: {}", gC_mnl.layout)
        debug("gC_raw layout: {}", gC_raw.layout)
        debug(
            "======================================================================\n\n"
        )

        #
        # Partition global tensor for TiledMMA_A/B/C
        #
        thr_mma = tiled_mma.get_slice(0)
        thr_mma_sfb = tiled_mma_sfb.get_slice(0)
        # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
        tCgA = thr_mma.partition_A(gA_mkl)
        # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
        tCgB = thr_mma.partition_B(gB_nkl)
        # (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
        tCgSFA = thr_mma.partition_A(gSFA_mkl)
        # (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
        tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
        # (MMA, MMA_M, MMA_N, RestM, RestN, RestL)
        tCgC = thr_mma.partition_C(gC_mnl)
        tCgC_raw = thr_mma.partition_C(gC_raw)

        debug("\n===================== TiledMMA Partitions =====================")
        debug("tCgA layout: {}", tCgA.layout)
        debug("tCgB layout: {}", tCgB.layout)
        debug("tCgSFA layout: {}", tCgSFA.layout)
        debug("tCgSFB layout: {}", tCgSFB.layout)
        debug("tCgC layout: {}", tCgC.layout)
        debug("tCgC_raw layout:\n\t{}\n\t{}", tCgC_raw.layout, tCgC_raw.iterator)
        debug(f"\ntiled_mma:\n{tiled_mma}\n")
        debug(f"\ntiled_mma_sfb:\n{tiled_mma_sfb}\n")
        debug("===============================================================\n\n")

        #
        # Partition global/shared tensor for TMA load A/B
        #
        # TMA load A partition_S/D
        a_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestM, RestK, RestL)
        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 load B partition_S/D
        b_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestN, RestK, RestL)
        # Note: For B, there's no M clustering (cluster_shape_mn[0]=1), so all CTAs
        # independently load their own B tiles. Use coord 0 since b_cta_layout has size 1.
        tBsB, tBgB = cpasync.tma_partition(
            tma_atom_b,
            0,
            b_cta_layout,
            cute.group_modes(sB, 0, 3),
            cute.group_modes(tCgB, 0, 3),
        )

        #  TMA load SFA partition_S/D
        sfa_cta_layout = a_cta_layout
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestM, RestK, RestL)
        tAsSFA, tAgSFA = cute.nvgpu.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 load SFB partition_S/D
        sfb_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
        )
        # ((atom_v, rest_v), STAGE)
        # ((atom_v, rest_v), RestN, RestK, RestL)
        # Note: For SFB, there's no M clustering (cluster_shape_mn[0]=1), so all CTAs
        # independently load their own SFB tiles. Use coord 0 since sfb_cta_layout has size 1.
        tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition(
            tma_atom_sfb,
            0,
            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)

        debug("\n===================== TMA Partitions =====================")
        debug("tAsA layout: {}", tAsA.layout)
        debug("tAgA layout: {}", tAgA.layout)
        debug("tBsB layout: {}", tBsB.layout)
        debug("tBgB layout: {}", tBgB.layout)
        debug("tAsSFA layout: {}", tAsSFA.layout)
        debug("tAgSFA layout: {}", tAgSFA.layout)
        debug("tBsSFB layout: {}", tBsSFB.layout)
        debug("tBgSFB layout: {}", tBgSFB.layout)
        debug("==========================================================\n\n")

        #
        # Partition shared/tensor memory tensor for TiledMMA_A/B/C
        #
        # (MMA, MMA_M, MMA_K, STAGE)
        tCrA = tiled_mma.make_fragment_A(sA)
        # (MMA, MMA_N, MMA_K, STAGE)
        tCrB = tiled_mma.make_fragment_B(sB)
        # (MMA, MMA_M, MMA_N)
        acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])

        # (MMA, MMA_M, MMA_N, STAGE)
        # not an actual tensor, just used for making layouts
        tCtAcc_fake = tiled_mma.make_fragment_C(
            cute.append(acc_shape, self.num_acc_stage)
        )

        debug("\n===================== MMA Fragments =====================")
        debug("tCrA layout: {}", tCrA.layout)
        debug("tCrB layout: {}", tCrB.layout)
        debug("acc_shape: {}", acc_shape)
        debug("tCtAcc_fake layout: {}", tCtAcc_fake.layout)
        debug("========================================================\n\n")

        #
        # Cluster wait before tensor memory alloc
        #
        pipeline_init_wait(cluster_shape_mn=self.cluster_shape_kn)

        #
        # Specialized TMA load warp
        #
        if warp_idx == self.tma_warp_id:
            self.load_warp(
                tiled_mma,
                tile_sched_params,
                ab_pipeline,
                tma_atom_a,
                tma_atom_b,
                tma_atom_sfa,
                tma_atom_sfb,
                tAgA,
                tBgB,
                tAgSFA,
                tBgSFB,
                tAsA,
                tBsB,
                tAsSFA,
                tBsSFB,
                a_full_mcast_mask,
                sfa_full_mcast_mask,
                sfb_full_mcast_mask,
            )

        #
        # Specialized MMA warp
        #
        if warp_idx == self.mma_warp_id:
            self.mma_warp(
                tiled_mma,
                tile_sched_params,
                tmem,
                tCtAcc_fake,
                sfa_smem_layout_staged,
                sfb_smem_layout_staged,
                sSFA,
                sSFB,
                tCrA,
                tCrB,
                ab_pipeline,
                acc_pipeline,
            )
        #
        # Specialized epilogue warps
        #
        if warp_idx < self.mma_warp_id:
            if cutlass.const_expr(self.ksplits == 1):
                self.epi_warp(
                    warp_idx,
                    tidx,
                    tiled_mma,
                    tile_sched_params,
                    tmem,
                    tCtAcc_fake,
                    tCgC,
                    epi_tile,
                    sC,
                    tma_atom_c,
                    acc_pipeline,
                )
            else:
                self.epi_warp_ksplit(
                    warp_idx,
                    tidx,
                    tiled_mma,
                    tile_sched_params,
                    tmem,
                    tCtAcc_fake,
                    tCgC_raw,
                    epi_tile,
                    sC,
                    red_pipeline,
                    acc_pipeline,
                )

    @cute.jit
    def load_warp(
        self,
        tiled_mma: cute.TiledMma,
        tile_sched_params: TileSchedulerParams,
        ab_pipeline: pipeline.PipelineTmaUmma,
        tma_atom_a: cute.CopyAtom,
        tma_atom_b: cute.CopyAtom,
        tma_atom_sfa: cute.CopyAtom,
        tma_atom_sfb: cute.CopyAtom,
        tAgA: cute.Tensor,
        tBgB: cute.Tensor,
        tAgSFA: cute.Tensor,
        tBgSFB: cute.Tensor,
        tAsA: cute.Tensor,
        tBsB: cute.Tensor,
        tAsSFA: cute.Tensor,
        tBsSFB: cute.Tensor,
        a_full_mcast_mask,
        sfa_full_mcast_mask,
        sfb_full_mcast_mask,
    ):
        """
        Specialized TMA load warp for loading A/B/SFA/SFB tiles.
        """
        #
        # Persistent tile scheduling loop
        #
        tile_sched = create_tile_scheduler(
            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
        )

        tAgA = self.tma_gmem_ksplit_partition(tAgA)
        tBgB = self.tma_gmem_ksplit_partition(tBgB)
        tAgSFA = self.tma_gmem_ksplit_partition(tAgSFA)
        tBgSFB = self.tma_gmem_ksplit_partition(tBgSFB)

        tma_debug(
            "\n===================== Load Warp KSPLIT Tensors =====================\n"
            "tAgA:\n\t{}\n\t{}\n"
            "tBgB:\n\t{}\n\t{}\n"
            "tAgSFA:\n\t{}\n\t{}\n"
            "tBgSFB:\n\t{}\n\t{}\n"
            "====================================================================\n\n",
            tAgA.layout,
            tAgA.iterator,
            tBgB.layout,
            tBgB.iterator,
            tAgSFA.layout,
            tAgSFA.iterator,
            tBgSFB.layout,
            tBgSFB.iterator,
        )

        while self.is_work_tile_valid(work_tile):
            # Get tile coord from tile scheduler
            m_tile, n_tile, ksplit_tile, l_tile = get_mnkl_indices(work_tile)

            tma_debug(
                "Load warp: block_idx=({},{},{}), m_tile={}, n_tile={}, ksplit_tile={}\n",
                cute.arch.block_idx()[0],
                cute.arch.block_idx()[1],
                cute.arch.block_idx()[2],
                m_tile,
                n_tile,
                ksplit_tile,
            )

            #
            # Slice to per mma tile index
            #
            # ((atom_v, rest_v), RestK)
            tAgA_slice = tAgA[(None, m_tile, (None, ksplit_tile), l_tile)]
            # ((atom_v, rest_v), RestK)
            tBgB_slice = tBgB[(None, n_tile, (None, ksplit_tile), l_tile)]

            # ((atom_v, rest_v), RestK)
            tAgSFA_slice = tAgSFA[(None, m_tile, (None, ksplit_tile), l_tile)]

            slice_n = n_tile
            if cutlass.const_expr(self.mma_tiler[1] == 64):
                slice_n = n_tile // 2
            # ((atom_v, rest_v), RestK)
            tBgSFB_slice = tBgSFB[(None, slice_n, (None, ksplit_tile), l_tile)]

            # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
            ab_producer_state.reset_count()
            peek_ab_empty_status = cutlass.Boolean(1)
            if ab_producer_state.count < self.ksplits_tile_cnt:
                peek_ab_empty_status = ab_pipeline.producer_try_acquire(
                    ab_producer_state
                )
            #
            # Tma load loop
            #
            for k_tile in cutlass.range(0, self.ksplits_tile_cnt, 1, unroll=1):
                # Conditionally wait for AB buffer empty
                ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)

                smem_stage = ab_producer_state.index
                tma_mbar_ptr = ab_pipeline.producer_get_barrier(ab_producer_state)

                # TMA load A/B/SFA/SFB
                cute.copy(
                    tma_atom_a,
                    tAgA_slice[(None, ab_producer_state.count)],
                    tAsA[(None, smem_stage)],
                    tma_bar_ptr=tma_mbar_ptr,
                    mcast_mask=a_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_b,
                    tBgB_slice[(None, ab_producer_state.count)],
                    tBsB[(None, smem_stage)],
                    tma_bar_ptr=tma_mbar_ptr,
                )
                cute.copy(
                    tma_atom_sfa,
                    tAgSFA_slice[(None, ab_producer_state.count)],
                    tAsSFA[(None, smem_stage)],
                    tma_bar_ptr=tma_mbar_ptr,
                    mcast_mask=sfa_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_sfb,
                    tBgSFB_slice[(None, ab_producer_state.count)],
                    tBsSFB[(None, smem_stage)],
                    tma_bar_ptr=tma_mbar_ptr,
                    mcast_mask=sfb_full_mcast_mask,
                )

                # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
                ab_producer_state.advance()
                peek_ab_empty_status = cutlass.Boolean(1)
                if ab_producer_state.count < self.ksplits_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()

        #
        # Wait A/B buffer empty
        #
        ab_pipeline.producer_tail(ab_producer_state)

    @cute.jit
    def mma_warp(
        self,
        tiled_mma: cute.TiledMma,
        tile_sched_params: TileSchedulerParams,
        tmem: utils.TmemAllocator,
        tCtAcc_fake: cute.Tensor,
        sfa_smem_layout_staged: cute.Layout,
        sfb_smem_layout_staged: cute.Layout,
        sSFA: cute.Tensor,
        sSFB: cute.Tensor,
        tCrA: cute.Tensor,
        tCrB: cute.Tensor,
        ab_pipeline: pipeline.PipelineTmaUmma,
        acc_pipeline: pipeline.PipelineUmmaAsync,
    ):
        """
        Specialized MMA warp for performing matrix multiply-accumulate operations.
        """
        #
        # Bar sync for retrieve tensor memory ptr from shared mem
        #
        tmem.wait_for_alloc()

        #
        # Retrieving tensor memory ptr and make accumulator/SFA/SFB tensor
        #
        acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
        # Make accumulator tmem tensor
        # (MMA, MMA_M, MMA_N, STAGE)
        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,
        )
        # (MMA, MMA_M, MMA_K)
        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,
        )
        # (MMA, MMA_N, MMA_K)
        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)
        #
        # Partition for S2T copy of SFA/SFB
        #
        (
            tiled_copy_s2t_sfa,
            tCsSFA_compact_s2t,
            tCtSFA_compact_s2t,
        ) = self.mma_sf_s2t_copy(sSFA, tCtSFA)
        (
            tiled_copy_s2t_sfb,
            tCsSFB_compact_s2t,
            tCtSFB_compact_s2t,
        ) = self.mma_sf_s2t_copy(sSFB, tCtSFB)

        # mma_debug(
        #     f"\n===================== MMA Warp TMEM Tensors =====================\n"
        #     f"tCtAcc_base:\n\t{tCtAcc_base.layout}\n\t{tCtAcc_base.iterator}\n"
        #     f"tCtSFA:\n\t{tCtSFA.layout}\n\t{tCtSFA.iterator}\n"
        #     f"tCtSFB:\n\t{tCtSFB.layout}\n\t{tCtSFB.iterator}\n"
        #     f"================================================================\n\n"
        # )

        # mma_debug(
        #     f"\n===================== S2T Copy Partitions =====================\n"
        #     f"tCsSFA_compact_s2t:\n\t{tCsSFA_compact_s2t.layout}\n\t{tCsSFA_compact_s2t.iterator}\n"
        #     f"tCtSFA_compact_s2t:\n\t{tCtSFA_compact_s2t.layout}\n\t{tCtSFA_compact_s2t.iterator}\n"
        #     f"tCsSFB_compact_s2t:\n\t{tCsSFB_compact_s2t.layout}\n\t{tCsSFB_compact_s2t.iterator}\n"
        #     f"tCtSFB_compact_s2t:\n\t{tCtSFB_compact_s2t.layout}\n\t{tCtSFB_compact_s2t.iterator}\n"
        #     f"\ntiled_copy_s2t_sfa:\n{tiled_copy_s2t_sfa}\n"
        #     f"\ntiled_copy_s2t_sfb:\n{tiled_copy_s2t_sfb}\n"
        #     f"===============================================================\n\n"
        # )

        #
        # Persistent tile scheduling loop
        #
        tile_sched = create_tile_scheduler(
            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 self.is_work_tile_valid(work_tile):
            # Get tile coord from tile scheduler
            mnkl_indices = get_mnkl_indices(work_tile)
            mma_tile_coord_mnl = (
                mnkl_indices[0],
                mnkl_indices[1],
                mnkl_indices[3],
            )

            # Set tensor memory buffer for current tile
            # (MMA, MMA_M, MMA_N)
            tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]

            # Peek (try_wait) AB buffer full for k_tile = 0
            ab_consumer_state.reset_count()
            peek_ab_full_status = cutlass.Boolean(1)
            if ab_consumer_state.count < self.ksplits_tile_cnt:
                peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)

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

            tCtSFB_mma = tCtSFB

            if cutlass.const_expr(self.mma_tiler[1] == 64):
                # Move in increments of 64 columns of SFB
                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 the ACCUMULATE field for each tile
            #
            tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

            #
            # Mma mainloop
            #
            for k_tile in range(self.ksplits_tile_cnt):
                smem_stage = ab_consumer_state.index
                # Conditionally 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,
                    smem_stage,
                )
                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,
                )

                # tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB
                num_kblocks = cute.size(tCrA, mode=[2])
                # assert num_kblocks == mma_inst_tile_k
                for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
                    kblock_coord = (
                        None,
                        None,
                        kblock_idx,
                        smem_stage,
                    )

                    # Set SFA/SFB tensor to tiled_mma
                    sf_kblock_coord = (None, None, kblock_idx)
                    # this updates the smem descriptor for sfa and sfb
                    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,
                    )

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

                # Async arrive AB buffer empty
                ab_pipeline.consumer_release(ab_consumer_state)

                # Peek (try_wait) AB buffer full for k_tile = k_tile + 1
                ab_consumer_state.advance()
                peek_ab_full_status = cutlass.Boolean(1)
                if ab_consumer_state.count < self.ksplits_tile_cnt:
                    peek_ab_full_status = ab_pipeline.consumer_try_wait(
                        ab_consumer_state
                    )

            #
            # Async arrive accumulator buffer full
            #
            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()

        #
        # Wait for accumulator buffer empty
        #
        acc_pipeline.producer_tail(acc_producer_state)

    @cute.jit
    def epi_warp_ksplit(
        self,
        warp_idx: cutlass.Int32,
        tidx: cutlass.Int32,
        tiled_mma: cute.TiledMma,
        tile_sched_params: TileSchedulerParams,
        tmem: utils.TmemAllocator,
        tCtAcc_fake: cute.Tensor,
        tCgC: cute.Tensor,
        epi_tile: cute.Tile,
        sAccRed: cute.Tensor,
        red_pipeline: pipeline.PipelineTmaAsync,
        acc_pipeline: pipeline.PipelineUmmaAsync,
    ):
        """
        Specialized epilogue warps for storing results to global memory with k-split reduction.

        Each CTA handles a portion of the N-tile and receives partials from peer CTAs
        via st.async.shared::cluster, sums in FP32, then stores directly to global memory.
        """
        #
        # Alloc tensor memory buffer
        #
        tmem.allocate(self.num_tmem_alloc_cols)

        #
        # Bar sync for retrieve tensor memory ptr from shared memory
        #
        tmem.wait_for_alloc()

        #
        # Retrieving tensor memory ptr and make accumulator tensor
        #
        acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
        # (MMA, MMA_M, MMA_N, STAGE)
        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,
        )
        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)

        tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, cute.Float32)
        tiled_copy_r2s, tRS_rC_send, tRS_sC = self.epilog_smem_copy_and_partition(
            tiled_copy_t2r, tTR_rC, epi_tidx, sAccRed
        )
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)

        #
        # Persistent tile scheduling loop
        #
        tile_sched = create_tile_scheduler(
            tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
        )
        work_tile = tile_sched.initial_work_tile_info()
        my_ksplit_idx = tile_sched.k_split_idx

        # (KSPLIT_TILE_M, KSPLIT_TILE_N, KSPLIT_M, KSPLIT_N, RestM, RestN, RestL)
        gC_ksplit = cute.flat_divide(
            tCgC[((None, None), 0, 0, None, None, None)], epi_tile
        )
        gC_ksplit = gC_ksplit[(None, None, None, my_ksplit_idx, None, None, None)]
        # (CPY, CPY_M, CPY_N, KSPLIT_M, RestM, RestN, RestL)
        bRG_gC = thr_copy_t2r.partition_D(gC_ksplit)

        debug(
            "\n===================== Epilogue Warp Tensors =====================\n"
            "tCtAcc_base:\n\t{}\n\t{}\n"
            "tTR_tAcc_base:\n\t{}\n\t{}\n"
            "tTR_rAcc:\n\t{}\n\t{}\n"
            "tTR_rC:\n\t{}\n\t{}\n"
            "tRS_rC:\n\t{}\n\t{}\n"
            "tRS_sC:\n\t{}\n\t{}\n"
            "gC_ksplit:\n\t(KSPLIT_TILE_M, KSPLIT_TILE_N, KSPLIT_M, KSPLIT_N, RestM, RestN, RestL)\n\t{}\n\t{}\n"
            "bRG_gC:\n\t(CPY, CPY_M, CPY_N, KSPLIT_M, RestM, RestN, RestL)\n\t{}\n\t{}\n"
            "epi_tile:\n{}\n"
            f"\ntiled_copy_t2r:\n{tiled_copy_t2r}\n"
            f"\nthr_copy_t2r:\n{thr_copy_t2r}\n"
            f"\ntiled_copy_r2s:\n{tiled_copy_r2s}\n"
            f"\nthr_copy_r2s:\n{thr_copy_r2s}\n"
            "================================================================\n\n",
            tCtAcc_base.layout,
            tCtAcc_base.iterator,
            tTR_tAcc_base.layout,
            tTR_tAcc_base.iterator,
            tTR_rAcc.layout,
            tTR_rAcc.iterator,
            tTR_rC.layout,
            tTR_rC.iterator,
            tRS_rC_send.layout,
            tRS_rC_send.iterator,
            tRS_sC.layout,
            tRS_sC.iterator,
            gC_ksplit.layout,
            gC_ksplit.iterator,
            bRG_gC.layout,
            bRG_gC.iterator,
            epi_tile,
        )

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

        # Threads/warps participating in tma store pipeline
        c_producer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread,
            self.threads_per_warp * len(self.epilog_warp_id),
        )
        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=self.num_c_stage,
            producer_group=c_producer_group,
        )

        red_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, self.num_acc_stage
        )
        red_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, self.num_acc_stage
        )

        cluster_n_base_rank = cute.arch.block_in_cluster_idx()[1] * self.ksplits

        copy_acc_red_atom = cute.make_copy_atom(
            cute.nvgpu.CopyUniversalOp(), self.acc_dtype
        )
        copy_d_atom = cute.make_copy_atom(
            cute.nvgpu.CopyUniversalOp(), self.d_dtype, num_bits_per_copy=256
        )
        tSR_rC_local = cute.make_rmem_tensor_like(tRS_rC_send, dtype=self.acc_dtype)
        tSR_rC_recv = cute.make_rmem_tensor_like(tSR_rC_local, self.acc_dtype)
        tDrD = cute.make_rmem_tensor_like(tSR_rC_local, self.d_dtype)

        while self.is_work_tile_valid(work_tile):
            # Get tile coord from tile scheduler
            mnkl_indices = get_mnkl_indices(work_tile)

            bRG_gC_slice = bRG_gC[(None, 0, 0, 0, 0, mnkl_indices[1], mnkl_indices[3])]

            acc_stage = acc_consumer_state.index

            # Set tensor memory buffer for current tile
            # (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
            tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_stage)]

            #
            # Wait for accumulator buffer full
            #
            # Wait for peer remotes to finish reading partials
            if warp_idx == self.epilog_warp_id[0]:
                red_pipeline.producer_acquire(red_producer_state)

            acc_pipeline.consumer_wait(acc_consumer_state)

            # ksplit_peer is the CTA rank for remote communication
            ksplit_peer = ((my_ksplit_idx + 1) % self.ksplits) + cluster_n_base_rank
            # peer_ksplit_idx is the k-split index of the peer (for local tensor indexing)
            peer_ksplit_idx = (my_ksplit_idx + 1) % self.ksplits

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

            # Index shared memory by k-split index (0 or 1), not CTA rank
            tRS_sC_send = tRS_sC[(None, None, None, my_ksplit_idx)]
            tSR_sC_recv = tRS_sC[(None, None, None, peer_ksplit_idx)]

            # Each ksplit peer CTA copies
            # 1. to local smem[my_ksplit_idx]
            # 2. from local smem[my_ksplit_idx] -> dsmem[my_ksplit_idx]
            # To support more than 2 k-splits, we'll need buffer space in smem for each k split
            #  - or we can use red.shared::cluster

            # Index tensor memory by k-split index, not CTA rank
            tTR_tAcc_mn = tTR_tAcc[(None, None, None, peer_ksplit_idx)]
            cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

            acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load().to(self.c_dtype)
            tRS_rC_send.store(acc_vec)

            cute.copy(
                tiled_copy_r2s,
                tRS_rC_send,
                dst=tRS_sC_send,
            )
            # Fence and barrier to make sure shared memory store is visible to TMA store
            cute.arch.fence_proxy(
                cute.arch.ProxyKind.async_shared,
                space=cute.arch.SharedSpace.shared_cta,
            )

            with cute.arch.elect_one():
                cp_async_bulk_remote(
                    # tRS_sC_remote.iterator,
                    tRS_sC_send.iterator,
                    tRS_sC_send.iterator,
                    red_pipeline.sync_object_full.get_barrier(red_producer_state.index),
                    self.red_send_bytes_per_warp,
                    ksplit_peer,
                )

            red_producer_state.advance()

            cute.copy(
                tiled_copy_t2r, tTR_tAcc[(None, None, None, my_ksplit_idx)], tTR_rAcc
            )
            acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load().to(self.c_dtype)
            tSR_rC_local.store(acc_vec)

            red_pipeline.consumer_wait(red_consumer_state)

            # TODO: subtile dis shit, like in fp_gemm_1.py

            tSR_rC_red_accum = tSR_rC_local.load()
            cute.copy(copy_acc_red_atom, tSR_sC_recv, tSR_rC_recv)
            tSR_rC_red_accum += tSR_rC_recv.load()
            tDrD.store(tSR_rC_red_accum.to(self.d_dtype))
            # alwaysprint(tDrD.layout)
            # alwaysprint(bRG_gC_slice.layout)
            # cute.copy(copy_d_atom, tDrD[((0, None), 0, 0)], bRG_gC_slice[((None, 0),)])
            cute.autovec_copy(tDrD, bRG_gC_slice)

            # Release the reduction pipeline buffers
            # select threads are predesignated to release. no need to guard with elect_one.
            red_pipeline.consumer_release(red_consumer_state)

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

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

        #
        # Dealloc the tensor memory buffer
        #
        tmem.relinquish_alloc_permit()
        self.epilog_sync_barrier.arrive_and_wait()
        tmem.free(acc_tmem_ptr)
        #
        # Wait for C store complete
        #
        c_pipeline.producer_tail()

    def mma_sf_s2t_copy(
        self,
        sSF: cute.Tensor,
        tSF: cute.Tensor,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        # (MMA, MMA_MN, MMA_K, STAGE)
        tCsSF_compact = cute.filter_zeros(sSF)
        # (MMA, MMA_MN, MMA_K)
        tCtSF_compact = cute.filter_zeros(tSF)

        # Make S2T CopyAtom and tiledCopy
        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)

        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
        tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
        tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t, tCsSF_compact_s2t_
        )
        # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
        tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)

        return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t

    def epilog_tmem_copy_and_partition(
        self,
        tidx: cutlass.Int32,
        tAcc: cute.Tensor,
        gC_mnl: cute.Tensor,
        epi_tile: cute.Tile,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        # Make tiledCopy for tensor memory load
        copy_atom_t2r = sm100_utils.get_tmem_load_op(
            self.mma_tiler,
            self.cd_majorness,
            self.c_dtype,
            self.acc_dtype,
            epi_tile,
            False,
        )
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE)
        tAcc_epi = cute.flat_divide(
            tAcc[((None, None), 0, 0, None)],
            epi_tile,
        )
        # (EPI_TILE_M, EPI_TILE_N)
        tiled_copy_t2r = tcgen05.make_tmem_copy(
            copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
        )

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

        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_mnl_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )
        # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
        # (T2R, T2R_M, T2R_N)
        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: cute.TiledCopy,
        tTR_rC: cute.Tensor,
        tidx: cutlass.Int32,
        sC: cute.Tensor,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        copy_atom_r2s = sm100_utils.get_smem_store_op(
            self.cd_majorness, self.c_dtype, self.acc_dtype, tiled_copy_t2r
        )
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
        # (R2S, R2S_M, R2S_N, PIPE_D)
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)
        # (R2S, R2S_M, R2S_N)
        tRS_rC = tiled_copy_r2s.retile(tTR_rC)
        return tiled_copy_r2s, tRS_rC, tRS_sC

    def epilog_gmem_copy_and_partition(
        self,
        tidx: cutlass.Int32,
        atom: Union[cute.CopyAtom, cute.TiledCopy],
        gC_mnl: cute.Tensor,
        epi_tile: cute.Tile,
        sC: cute.Tensor,
    ) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
        # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
        gC_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )

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

    def tma_gmem_ksplit_partition(self, gT: cute.Tensor) -> cute.Tensor:
        return cute.logical_divide(gT, (None, None, self.ksplits_tile_cnt, None))

    def _get_k_splits(self, a_tensor: cute.Tensor):
        K = cute.size(a_tensor, mode=[1])
        ksplits = get_k_splits(K)

        ksplit_tiles = K // (self.mma_tiler[2] * ksplits)
        return ksplits, ksplit_tiles

    @staticmethod
    def _compute_stages(
        tiled_mma: cute.TiledMma,
        mma_tiler_mnk: Tuple[int, int, int],
        a_dtype: Type[cutlass.Numeric],
        b_dtype: Type[cutlass.Numeric],
        epi_tile: cute.Tile,
        c_dtype: Type[cutlass.Numeric],
        c_layout: utils.LayoutEnum,
        sf_dtype: Type[cutlass.Numeric],
        sf_vec_size: int,
        smem_capacity: int,
    ) -> Tuple[int, int, int]:
        # ACC stages
        num_acc_stage = 2
        num_c_stage = 2

        # Calculate smem layout and size for one stage of A, B, SFA, SFB and C
        a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
            tiled_mma,
            mma_tiler_mnk,
            a_dtype,
            1,  # a tmp 1 stage is provided
        )
        b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
            tiled_mma,
            mma_tiler_mnk,
            b_dtype,
            1,  # a tmp 1 stage is provided
        )
        sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,  # a tmp 1 stage is provided
        )
        sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,  # a tmp 1 stage is provided
        )

        c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(
            c_dtype,
            c_layout,
            epi_tile,
            1,
        )

        ab_bytes_per_stage = (
            cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
            + cute.size_in_bytes(b_dtype, b_smem_layout_staged_one)
            + cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
            + cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
        )
        mbar_helpers_bytes = 1024
        c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one)
        c_bytes = c_bytes_per_stage * num_c_stage

        # Calculate A/B/SFA/SFB stages:
        # Start with total smem per CTA (capacity / occupancy)
        # Subtract reserved bytes and initial C stages bytes
        # Divide remaining by bytes needed per A/B/SFA/SFB stage
        num_ab_stage = (
            smem_capacity - (mbar_helpers_bytes + c_bytes)
        ) // ab_bytes_per_stage

        # Refine epilogue stages:
        # Calculate remaining smem after allocating for A/B/SFA/SFB stages and reserved bytes
        # Add remaining unused smem to epilogue
        num_c_stage += (
            smem_capacity
            - ab_bytes_per_stage * num_ab_stage
            - (mbar_helpers_bytes + c_bytes)
        ) // c_bytes_per_stage

        return num_acc_stage, num_ab_stage, num_c_stage

    @staticmethod
    def _compute_grid(
        N_tiles: cutlass.Constexpr[cute.Int32],
        cluster_shape_kn: Tuple[int, int],
        max_active_clusters: cutlass.Constexpr,
    ) -> Tuple[TileSchedulerParams, Tuple[int, int, int]]:
        ksplits = cluster_shape_kn[0]

        if USE_SIMPLE_SCHEDULER:
            grid = SimpleTileSchedulerParams.get_grid_shape(
                cluster_shape_kn, N_tiles, max_active_clusters
            )
            n_stride = grid[1]  # num_persistent_n_ctas
            tile_sched_params = SimpleTileSchedulerParams(ksplits, N_tiles, n_stride)
        else:
            # For static scheduler with k-splits:
            # - Interpret "M" dimension as k_split, "N" dimension as n_tile
            # - problem_shape_ntile_mnl = (ksplits, N_tiles, 1)
            # - cluster_shape_mnk uses cluster_shape_kn (k_cluster, n_cluster, 1)
            #   where cluster_shape_kn[0] = k_splits, cluster_shape_kn[1] = n_cluster
            # - cluster_shape_mn is a "fake" shape only used for TMA multicasting
            problem_shape_ntile_mnl = (ksplits, N_tiles, 1)
            # Use cluster_shape_kn for the actual clustering
            cluster_shape_mnk = (cluster_shape_kn[0], cluster_shape_kn[1], 1)
            tile_sched_params = PersistentTileSchedulerParams(
                problem_shape_ntile_mnl,
                cluster_shape_mnk,
            )
            grid = StaticPersistentTileScheduler.get_grid_shape(
                tile_sched_params, max_active_clusters
            )

        hdebug("\n===================== Grid Configuration =====================")
        hdebug("N_tiles: {}", N_tiles)
        hdebug("k_splits: {}", ksplits)
        hdebug("cluster_shape_kn: {}", cluster_shape_kn)
        hdebug("grid: (k_splits={}, n_persistent={}, 1)", grid[0], grid[1])
        hdebug("==============================================================\n\n")

        # cute.printf("\n===================== Grid Configuration =====================")
        # cute.printf("N_tiles: {}", N_tiles)
        # cute.printf("k_splits: {}", ksplits)
        # cute.printf("cluster_shape_kn: {}", cluster_shape_kn)
        # cute.printf("grid: (k_splits={}, n_persistent={}, 1)", grid[0], grid[1])
        # cute.printf(
        #     "==============================================================\n\n"
        # )

        return tile_sched_params, grid

    @cute.jit
    def epi_warp(
        self,
        warp_idx: cutlass.Int32,
        tidx: cutlass.Int32,
        tiled_mma: cute.TiledMma,
        tile_sched_params: TileSchedulerParams,
        tmem: utils.TmemAllocator,
        tCtAcc_fake: cute.Tensor,
        tCgC: cute.Tensor,
        epi_tile: cute.Tile,
        sC: cute.Tensor,
        tma_atom_c: cute.CopyAtom,
        acc_pipeline: pipeline.PipelineUmmaAsync,
    ):
        """
        Specialized epilogue warps for storing results to global memory.
        """
        #
        # Alloc tensor memory buffer
        #
        tmem.allocate(self.num_tmem_alloc_cols)

        #
        # Bar sync for retrieve tensor memory ptr from shared memory
        #
        tmem.wait_for_alloc()

        #
        # Retrieving tensor memory ptr and make accumulator tensor
        #
        acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
        # (MMA, MMA_M, MMA_N, STAGE)
        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,
        )

        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(
            epi_tidx, tma_atom_c, tCgC, epi_tile, sC
        )

        debug(
            f"\n===================== Epilogue Warp Tensors =====================\n"
            f"tCtAcc_base:\n\t{tCtAcc_base.layout}\n\t{tCtAcc_base.iterator}\n"
            f"tTR_tAcc_base:\n\t{tTR_tAcc_base.layout}\n\t{tTR_tAcc_base.iterator}\n"
            f"tTR_rAcc:\n\t{tTR_rAcc.layout}\n\t{tTR_rAcc.iterator}\n"
            f"tTR_rC:\n\t{tTR_rC.layout}\n\t{tTR_rC.iterator}\n"
            f"tRS_rC:\n\t{tRS_rC.layout}\n\t{tRS_rC.iterator}\n"
            f"tRS_sC:\n\t{tRS_sC.layout}\n\t{tRS_sC.iterator}\n"
            f"bSG_sC:\n\t{bSG_sC.layout}\n\t{bSG_sC.iterator}\n"
            f"bSG_gC_partitioned:\n\t{bSG_gC_partitioned.layout}\n\t{bSG_gC_partitioned.iterator}\n"
            f"\ntiled_copy_t2r:\n{tiled_copy_t2r}\n"
            f"\ntiled_copy_r2s:\n{tiled_copy_r2s}\n"
            f"================================================================\n\n"
        )

        #
        # Persistent tile scheduling loop
        #
        tile_sched = create_tile_scheduler(
            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
        )

        # Threads/warps participating in tma store pipeline
        c_producer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread,
            self.threads_per_warp * len(self.epilog_warp_id),
        )
        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=self.num_c_stage,
            producer_group=c_producer_group,
        )

        while self.is_work_tile_valid(work_tile):
            # Get tile coord from tile scheduler
            mnkl_indices = get_mnkl_indices(work_tile)
            mma_tile_coord_mnl = (
                mnkl_indices[0],
                mnkl_indices[1],
                mnkl_indices[3],
            )

            #
            # Slice to per mma tile index
            #
            # ((ATOM_V, REST_V), EPI_M, EPI_N)
            bSG_gC = bSG_gC_partitioned[
                (
                    None,
                    None,
                    None,
                    *mma_tile_coord_mnl,
                )
            ]

            # Set tensor memory buffer for current tile
            # (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
            tTR_tAcc = tTR_tAcc_base[
                (None, None, None, None, None, acc_consumer_state.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):
                # Load accumulator from tensor memory buffer to register
                #
                tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
                cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

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

                #
                # Store C to shared memory
                #
                c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
                cute.copy(
                    tiled_copy_r2s,
                    tRS_rC,
                    dst=tRS_sC[(None, None, None, c_buffer)],
                )
                # Fence and barrier to make sure shared memory store is visible to TMA store
                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 to global memory
                #
                if warp_idx == self.epilog_warp_id[0]:
                    cute.copy(
                        tma_atom_c,
                        bSG_sC[(None, c_buffer)],
                        bSG_gC[(None, subtile_idx)],
                    )
                    # Fence and barrier to make sure shared memory store is visible to TMA store
                    c_pipeline.producer_commit()
                    c_pipeline.producer_acquire()
                self.epilog_sync_barrier.arrive_and_wait()

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

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

        #
        # Dealloc the tensor memory buffer
        #
        tmem.relinquish_alloc_permit()
        self.epilog_sync_barrier.arrive_and_wait()
        tmem.free(acc_tmem_ptr)
        #
        # Wait for C store complete
        #
        c_pipeline.producer_tail()

    @cute.jit
    def epi_warp_ksplit_more_general(
        self,
        warp_idx: cutlass.Int32,
        tidx: cutlass.Int32,
        tiled_mma: cute.TiledMma,
        tile_sched_params: TileSchedulerParams,
        tmem: utils.TmemAllocator,
        tCtAcc_fake: cute.Tensor,
        tCgC: cute.Tensor,
        epi_tile: cute.Tile,
        sAccRed: cute.Tensor,
        red_pipeline: pipeline.PipelineTmaAsync,
        acc_pipeline: pipeline.PipelineUmmaAsync,
    ):
        """
        Specialized epilogue warps for storing results to global memory with k-split reduction.

        Each CTA handles a portion of the N-tile and receives partials from peer CTAs
        via st.async.shared::cluster, sums in FP32, then stores directly to global memory.
        """
        #
        # Alloc tensor memory buffer
        #
        tmem.allocate(self.num_tmem_alloc_cols)

        #
        # Bar sync for retrieve tensor memory ptr from shared memory
        #
        tmem.wait_for_alloc()

        #
        # Retrieving tensor memory ptr and make accumulator tensor
        #
        acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
        # (MMA, MMA_M, MMA_N, STAGE)
        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,
        )
        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)

        tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, cute.Float32)
        tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
            tiled_copy_t2r, tTR_rC, epi_tidx, sAccRed
        )
        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)

        #
        # Persistent tile scheduling loop
        #
        tile_sched = create_tile_scheduler(
            tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
        )
        work_tile = tile_sched.initial_work_tile_info()
        my_ksplit_idx = tile_sched.k_split_idx

        # (KSPLIT_TILE_M, KSPLIT_TILE_N, KSPLIT_M, KSPLIT_N, RestM, RestN, RestL)
        gC_ksplit = cute.flat_divide(
            tCgC[((None, None), 0, 0, None, None, None)], epi_tile
        )
        gC_ksplit = gC_ksplit[(None, None, None, my_ksplit_idx, None, None, None)]
        # (CPY, CPY_M, CPY_N, KSPLIT_M, RestM, RestN, RestL)
        bRG_gC = thr_copy_t2r.partition_D(gC_ksplit)

        debug(
            "\n===================== Epilogue Warp Tensors =====================\n"
            "tCtAcc_base:\n\t{}\n\t{}\n"
            "tTR_tAcc_base:\n\t{}\n\t{}\n"
            "tTR_rAcc:\n\t{}\n\t{}\n"
            "tTR_rC:\n\t{}\n\t{}\n"
            "tRS_rC:\n\t{}\n\t{}\n"
            "tRS_sC:\n\t{}\n\t{}\n"
            "gC_ksplit:\n\t(KSPLIT_TILE_M, KSPLIT_TILE_N, KSPLIT_M, KSPLIT_N, RestM, RestN, RestL)\n\t{}\n\t{}\n"
            "bRG_gC:\n\t(CPY, CPY_M, CPY_N, KSPLIT_M, RestM, RestN, RestL)\n\t{}\n\t{}\n"
            "epi_tile:\n{}\n"
            f"\ntiled_copy_t2r:\n{tiled_copy_t2r}\n"
            f"\nthr_copy_t2r:\n{thr_copy_t2r}\n"
            f"\ntiled_copy_r2s:\n{tiled_copy_r2s}\n"
            f"\nthr_copy_r2s:\n{thr_copy_r2s}\n"
            "================================================================\n\n",
            tCtAcc_base.layout,
            tCtAcc_base.iterator,
            tTR_tAcc_base.layout,
            tTR_tAcc_base.iterator,
            tTR_rAcc.layout,
            tTR_rAcc.iterator,
            tTR_rC.layout,
            tTR_rC.iterator,
            tRS_rC.layout,
            tRS_rC.iterator,
            tRS_sC.layout,
            tRS_sC.iterator,
            gC_ksplit.layout,
            gC_ksplit.iterator,
            bRG_gC.layout,
            bRG_gC.iterator,
            epi_tile,
        )

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

        # Threads/warps participating in tma store pipeline
        c_producer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread,
            self.threads_per_warp * len(self.epilog_warp_id),
        )
        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=self.num_c_stage,
            producer_group=c_producer_group,
        )

        red_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, self.num_acc_stage
        )
        red_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, self.num_acc_stage
        )

        cluster_n_base_rank = cute.arch.block_in_cluster_idx()[1] * self.ksplits

        while self.is_work_tile_valid(work_tile):
            # for _ in cutlass.range_constexpr(1):
            # Get tile coord from tile scheduler
            mnkl_indices = get_mnkl_indices(work_tile)

            bRG_gC_slice = bRG_gC[(None, 0, 0, 0, 0, mnkl_indices[1], mnkl_indices[3])]

            acc_stage = acc_consumer_state.index

            # Set tensor memory buffer for current tile
            # (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
            tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_stage)]

            #
            # Wait for accumulator buffer full
            #
            # Producer acquires the reduction buffer (waits for empty, sets tx count)
            if warp_idx == self.epilog_warp_id[0]:
                red_pipeline.producer_acquire(red_producer_state)

            acc_pipeline.consumer_wait(acc_consumer_state)

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

            tRS_sC_local = tRS_sC[(None, None, None, my_ksplit_idx)]
            tSR_rC_local = cute.make_rmem_tensor_like(tRS_rC, dtype=cute.Float32)

            for i in cutlass.range_constexpr(1, self.ksplits, 1):
                assert self.ksplits == 2
                """
                TODO: store to dsmem in my_split_idx
                partition w.r.t. atom

                tRDS_rC layout:
                        ((4,8),1,1):((1,4),0,0)
                        raw_ptr(0x0000fffc9dfffb00: f32, rmem, align<32>)
                tRDS_sC layout:
                        ((4,8),1,1,(1,3)):((1,4),0,0,(0,4096))
                        raw_ptr(0x0000000000000880: f32, smem, align<128>)
                st_atom_shape = (4,)
                tRDS_rC = cute.logical_divide(tRS_rC, st_atom_shape)
                tRDS_sC = cute.logical_divide(
                    tRS_sC[(None, None, None, cur_ksplit)], st_atom_shape
                )
                debug(
                    "tRDS_rC layout:\n\t{}\n\t{}",
                    tRDS_rC.layout,
                    tRDS_rC.iterator,
                )
                debug(
                    "tRDS_sC layout:\n\t{}\n\t{}",
                    tRDS_sC.layout,
                    tRDS_sC.iterator,
                )
                """

                # Each ksplit peer CTA copies
                # 1. to local smem[my_ksplit_idx]
                # 2. from local smem[my_ksplit_idx] -> dsmem[my_ksplit_idx]
                # To support more than 2 k-splits, we'll need buffer space in smem for each k split
                #  - or we can use red.shared::cluster

                ksplit_peer = ((my_ksplit_idx + i) % self.ksplits) + cluster_n_base_rank
                tTR_tAcc_mn = tTR_tAcc[(None, None, None, ksplit_peer)]
                cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

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

                cute.copy(
                    tiled_copy_r2s,
                    tRS_rC,
                    dst=tRS_sC_local,
                )
                # Fence and barrier to make sure shared memory store is visible to TMA store
                cute.arch.fence_proxy(
                    cute.arch.ProxyKind.async_shared,
                    space=cute.arch.SharedSpace.shared_cta,
                )

                with cute.arch.elect_one():
                    cp_async_bulk_remote(
                        # tRS_sC_remote.iterator,
                        tRS_sC_local.iterator,
                        tRS_sC_local.iterator,
                        red_pipeline.sync_object_full.get_barrier(
                            red_producer_state.index
                        ),
                        self.red_send_bytes_per_warp,
                        ksplit_peer,
                    )

            red_producer_state.advance()

            cute.copy(
                tiled_copy_t2r, tTR_tAcc[(None, None, None, my_ksplit_idx)], tTR_rAcc
            )
            acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load().to(self.c_dtype)
            tSR_rC_local.store(acc_vec)

            red_pipeline.consumer_wait(red_consumer_state)

            # TODO: subtile dis shit, like in fp_gemm_1.py
            copy_acc_red_atom = cute.make_copy_atom(
                cute.nvgpu.CopyUniversalOp(), cute.Float32
            )
            tSR_rC_buf = cute.make_rmem_tensor_like(tSR_rC_local, cute.Float32)

            tSR_rC_local_ssa = tSR_rC_local.load()

            for cur_ksplit in cutlass.range_constexpr(self.ksplits):
                if my_ksplit_idx != cur_ksplit:
                    tSR_sC_red = tRS_sC[(None, None, None, cur_ksplit)]
                    cute.copy(copy_acc_red_atom, tSR_sC_red, tSR_rC_buf)
                    tSR_rC_local_ssa += tSR_rC_buf.load()

            tDrD = cute.make_rmem_tensor_like(tSR_rC_local, self.d_dtype)
            tDrD.store(tSR_rC_local_ssa.to(self.d_dtype))

            cute.autovec_copy(tDrD, bRG_gC_slice)

            # Release the reduction pipeline buffers
            # select threads are predesignated to release. no need to guard with elect_one.
            red_pipeline.consumer_release(red_consumer_state)

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

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

        # j
        # Dealloc the tensor memory buffer
        #
        tmem.relinquish_alloc_permit()
        self.epilog_sync_barrier.arrive_and_wait()
        tmem.free(acc_tmem_ptr)
        #
        # Wait for C store complete
        #
        c_pipeline.producer_tail()


@cute.jit
def gemm_fp4_host(
    a_ptr: cute.Pointer,
    b_ptr: cute.Pointer,
    sfa_ptr: cute.Pointer,
    sfb_ptr: cute.Pointer,
    c_ptr: cute.Pointer,
    n: cutlass.Constexpr[cute.Int32],
    k: cutlass.Constexpr[cute.Int32],
    max_active_clusters: cutlass.Constexpr[cute.Int32],
):
    m = 128
    l = 1

    # Setup attributes that depend on gemm inputs
    a_tensor = cute.make_tensor(
        a_ptr,
        cute.make_layout(
            (m, k, l),
            stride=(k, 1, m * k),
        ),
    )
    b_tensor = cute.make_tensor(
        b_ptr,
        cute.make_layout(
            (n, k, l),
            stride=(k, 1, n * k),
        ),
    )
    c_tensor = cute.make_tensor(
        c_ptr, cute.make_layout((m, n, l), stride=(n, 1, m * n))
    )
    # Setup sfa/sfb tensor by filling A/B tensor to scale factor atom layout
    # ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL)
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

    # ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL)
    sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
    sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)

    FP4GEMMKernel()(
        a_tensor,
        b_tensor,
        sfa_tensor,
        sfb_tensor,
        c_tensor,
        max_active_clusters,
    )


# Global cache for compiled kernel
KERNEL_CACHE = {}


def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    # Please ensure rows and cols are multiples of 128 and 4 respectively
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()


def ref_kernel(
    data: input_t,
) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled GEMM.
    """
    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data

    # Get dimensions from MxNxL layout
    _, _, l = c_ref.shape

    # Call torch._scaled_mm to compute the GEMM result
    for l_idx in range(l):
        # Convert the scale factor tensor to blocked format
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])
        # (m, k) @ (n, k).T -> (m, n)
        res = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b.cuda(),
            bias=None,
            out_dtype=torch.float16,
        )
        c_ref[:, :, l_idx] = res
    return c_ref


def cute_compile(*args, **kwargs):
    options = []
    if COMPILE_PTX:
        cubin_option = cute.KeepCUBIN(True)
        # cubin_option.dump_path = DUMP_PATH
        ptx_option = cute.KeepPTX(True)
        # ptx_option.dump_path = DUMP_PATH
        options.extend([cubin_option, ptx_option])
        options.append(cute.GenerateLineInfo(True))

    kernel = cute.compile(*args, **kwargs)

    # Disassemble cubin if COMPILE_PTX is enabled
    import socket

    hostname = socket.gethostname()

    if COMPILE_PTX and hostname == "thor":
        dump_dir = Path(os.environ.get("CUTE_DSL_DUMP_DIR", "."))
        function_name = kernel.function_name
        base = dump_dir / f"{function_name}.{detect_gpu_arch('')}"

        cubin_path = Path(f"{base}.cubin")
        asm_path = Path(f"{base}.asm")

        if cubin_path.exists():
            try:
                with open(asm_path, "w") as f:
                    subprocess.Popen(
                        ["disasm", str(cubin_path)],
                        stdout=f,
                        stderr=subprocess.DEVNULL,
                    )
                # print(f"Disassembling cubin to {asm_path}")
            except FileNotFoundError:
                print("disasm not found. Please ensure CUDA toolkit is in PATH.")
        else:
            print(f"Cubin file not found at {cubin_path}")

    return kernel


# This function is used to compile the kernel once and cache it and then allow users to
# run the kernel multiple times to get more accurate timing results.
def _compile_kernel(n: cutlass.Constexpr[cute.Int32], k: cutlass.Constexpr[cute.Int32]):
    """
    Compile the kernel once and cache it.
    This should be called before any timing measurements.

    Returns:
        The compiled kernel function
    """
    if (n, k) in KERNEL_CACHE:
        return KERNEL_CACHE[(n, k)]

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    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)

    cutlass.cuda.initialize_cuda_context()
    k_splits = get_k_splits(k)
    cluster_shape_kn = (k_splits, cluster_shape_mn[1])
    max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters(
        cluster_shape_kn[0] * cluster_shape_kn[1]
    )

    args = (
        gemm_fp4_host,
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
        n,
        k,
        max_active_clusters,
    )

    KERNEL_CACHE[(n, k)] = cute.compile[cute.KeepCUBIN(), cute.KeepPTX()](*args)

    return KERNEL_CACHE[(n, k)]


BENCH_DIMS = [
    (7168, 16384),
    (4096, 7168),
    (7168, 2048),
]


def compile_kernel():
    """
    Pre-compile kernels for all test case dimensions.
    This ensures compilation time is not included in benchmark results.
    """
    # All (n, k) combinations from task.yml tests/benchmarks
    for n, k in BENCH_DIMS:
        _compile_kernel(n, k)


def custom_kernel(data: input_t) -> output_t:
    """
    Execute the block-scaled GEMM kernel.

    This is the main entry point called by the evaluation framework.
    It converts PyTorch tensors to CuTe tensors, launches the kernel,
    and returns the result.

    Args:
        data: Tuple of (a, b, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c) PyTorch tensors
            a: [m, k, l] - Input matrix in float4e2m1fn
            b: [n, k, l] - Input vector in float4e2m1fn
            sfa_ref: [m, k, l] - Scale factors in float8_e4m3fn, used by reference implementation
            sfb_ref: [n, k, l] - Scale factors in float8_e4m3fn, used by reference implementation
            sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn
            sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
            c: [m, n, l] - Output vector in float16

    Returns:
        Output tensor c with computed results
    """

    a, b, _, _, sfa_permuted, sfb_permuted, c = data

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

    if (n, k) not in BENCH_DIMS or m != 128 or l != 1:
        return ref_kernel(data)

    # Ensure kernel is compiled (will use cached version if available)
    # To avoid the compilation overhead, we compile the kernel once and cache it.
    compiled_func = _compile_kernel(n, k)

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    sfa_ptr = make_ptr(
        sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
    )
    sfb_ptr = make_ptr(
        sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
    )

    # Execute the compiled kernel
    compiled_func(
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
    )

    return c
scrolls · 3450 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON