Skip to content
KernelIndex
Search⌘K

submission 85147

shellsmile15795 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v15.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-85147?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 GEMVsuite of 3 cases
NVIDIA B200
23.1µs
#59 of 678
2025-11-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7012f0f506bc5a61a95e6b4bc6a430b63ed17f8ffdc589915736ec4868b547e8
license declaredunknown
license concludedunknown
authorsshellsmile15795
imported2026-08-15

Techniques

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

fused-epilogueself.epi_tile = sm100_utils.compute_epilogue_tile_shape(
mbarrierself.cta_sync_barrier = pipeline.NamedBarrier(
persistent-kernelclass PersistentTileSchedulerParams:
shared-memoryself.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
tcgen05tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission_v15.py1878 lines
import torch
from task import input_t, output_t

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils.blockscaled_layout as blockscaled_utils


from typing import Type, Tuple, Union

import cuda.bindings.driver as cuda
import torch

import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.torch as cutlass_torch
import cutlass.utils as utils
import cutlass.pipeline as pipeline
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import from_dlpack

from cutlass.cutlass_dsl import (
    Boolean,
    Integer,
    Int32,
    min,
    extract_mlir_values,
    new_from_mlir_values,
    dsl_user_op,
)
from cutlass._mlir import ir

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


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:
    """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

        # Cache values needed in the hot scheduler loop so we only emit simple
        # integer MADs per tile request.
        self._cluster_tile_multipliers = tuple(
            Int32(x) for x in (*self.params.cluster_shape_mn, 1)
        )
        self._num_problem_clusters = Int32(
            cute.size(self.params.problem_layout_ncluster_mnl)
        )

        self._problem_cluster_shape = None
        if self.params.swizzle_size == 1:
            layout_shape = self.params.problem_layout_ncluster_mnl.shape
            m_dim, n_dim, l_dim = layout_shape
            self._problem_cluster_shape = (
                Int32(cute.size(m_dim)),
                Int32(cute.size(n_dim)),
                Int32(cute.size(l_dim)),
            )

    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
        # Grid dimensions are sized to the cluster shape so we can use blockIdx
        # directly instead of paying for modulo instructions at runtime.
        cta_id_in_cluster = (
            Int32(bidx),
            Int32(bidy),
            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 < self._num_problem_clusters

        if self.params.swizzle_size == 1 and self._problem_cluster_shape is not None:
            linear_idx = current_work_linear_idx
            m_extent, n_extent, _ = self._problem_cluster_shape

            m_coord = linear_idx % m_extent
            tmp = linear_idx // m_extent
            n_coord = tmp % n_extent
            l_coord = tmp // n_extent
            cur_cluster_coord = (m_coord, n_coord, l_coord)
        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) * z + y
            for x, y, z in zip(
                cur_cluster_coord, self.cta_id_in_cluster, self._cluster_tile_multipliers
            )
        )

        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):
        advance = Int32(advance_count)
        self._current_work_linear_idx += advance * Int32(
            self.num_persistent_clusters
        )
        self._num_tiles_executed += advance

    @dsl_user_op
    def acquire_work_batch(self, *, batch_size: int = 1, loc=None, ip=None):
        if batch_size < 1:
            raise ValueError("batch_size must be >= 1")

        tiles: list[WorkTileInfo] = []
        base_linear_idx = self._current_work_linear_idx
        stride = Int32(self.num_persistent_clusters)

        for batch_idx in range(batch_size):
            tile_linear_idx = base_linear_idx + Int32(batch_idx) * stride
            tiles.append(
                self._get_current_work_for_linear_idx(
                    tile_linear_idx, loc=loc, ip=ip
                )
            )

        return tuple(tiles)

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

#
# NOTE: Optimized Kernel
#

class OptimizedKernel:
    def __init__(
        self,
        sf_vec_size: int,
        mma_tiler_mn: Tuple[int, int],
        cluster_shape_mn: Tuple[int, int],
    ):
        self.acc_dtype = cutlass.Float32
        self.sf_vec_size = sf_vec_size
        self.use_2cta_instrs = mma_tiler_mn[0] == 256
        self.cluster_shape_mn = cluster_shape_mn

        self.mma_tiler = (*mma_tiler_mn, 1)

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

        self.occupancy = 1

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

        self.cta_sync_barrier = pipeline.NamedBarrier(
            barrier_id=1,
            num_threads=self.threads_per_cta,
        )
        self.epilog_sync_barrier = pipeline.NamedBarrier(
            barrier_id=2,
            num_threads=32 * len(self.epilog_warp_id),
        )
        self.tmem_alloc_barrier = pipeline.NamedBarrier(
            barrier_id=3,
            num_threads=32 * 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

    def _setup_attributes(self):
        self.mma_inst_shape_mn = (
            self.mma_tiler[0],
            self.mma_tiler[1],
        )

        self.mma_inst_shape_mn_sfb = (
            self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
            cute.round_up(self.mma_inst_shape_mn[1], 128),
        )

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

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

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

        self.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.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2])
        self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1])
        self.num_mcast_ctas_sfb = cute.size(self.cluster_layout_sfb_vmnk.shape[1])
        self.is_a_mcast = self.num_mcast_ctas_a > 1
        self.is_b_mcast = self.num_mcast_ctas_b > 1
        self.is_sfb_mcast = self.num_mcast_ctas_sfb > 1

        self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
            self.cta_tile_shape_mnk,
            self.use_2cta_instrs,
            self.c_layout,
            self.c_dtype,
        )

        self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
            tiled_mma,
            self.mma_tiler,
            self.a_dtype,
            self.b_dtype,
            self.epi_tile,
            self.c_dtype,
            self.c_layout,
            self.sf_dtype,
            self.sf_vec_size,
            self.smem_capacity,
            self.occupancy,
        )

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

    @cute.jit
    def __call__(
        self,
        a_ptr: cute.Pointer,
        b_ptr: cute.Pointer,
        sfa_ptr: cute.Pointer,
        sfb_ptr: cute.Pointer,
        c_ptr: cute.Pointer,
        problem_size: tuple,
        max_active_clusters: cutlass.Constexpr,
        stream: cuda.CUstream,
        epilogue_op: cutlass.Constexpr = lambda x: x,
    ):
        m, n, k, l = problem_size
        m = cute.assume(m, divby=32)
        k = cute.assume(k, divby=32)

        a_layout = cute.make_ordered_layout((m, k, l), order=(1, 0, 2))
        a_tensor = cute.make_tensor(a_ptr, layout=a_layout)

        n_padded_128 = 128
        b_layout = cute.make_ordered_layout((n_padded_128, k, l), order=(1, 0, 2))
        b_tensor = cute.make_tensor(b_ptr, layout=b_layout)

        c_layout = cute.make_ordered_layout((m, n, l), order=(0, 1, 2))
        c_tensor = cute.make_tensor(c_ptr, layout=c_layout)

        sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
            a_tensor.shape, self.sf_vec_size
        )
        sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

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

        self.a_dtype: Type[cutlass.Numeric] = a_tensor.element_type
        self.b_dtype: Type[cutlass.Numeric] = b_tensor.element_type
        self.sf_dtype: Type[cutlass.Numeric] = sfa_tensor.element_type
        self.c_dtype: Type[cutlass.Numeric] = c_tensor.element_type
        self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor).mma_major_mode()
        self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor).mma_major_mode()
        self.c_layout = utils.LayoutEnum.from_tensor(c_tensor)

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

        self._setup_attributes()

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

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

        a_op = sm100_utils.cluster_shape_to_tma_atom_A(
            self.cluster_shape_mn, tiled_mma.thr_id
        )
        a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
        tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
            a_op,
            a_tensor,
            a_smem_layout,
            self.mma_tiler,
            tiled_mma,
            self.cluster_layout_vmnk.shape,
        )

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

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

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

        if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
            x = tma_tensor_sfb.stride[0][1]
            y = cute.ceil_div(tma_tensor_sfb.shape[0][1], 4)

            new_shape = (
                (tma_tensor_sfb.shape[0][0], ((2, 2), y)),
                tma_tensor_sfb.shape[1],
                tma_tensor_sfb.shape[2],
            )

            x_times_3 = 3 * x
            new_stride = (
                (tma_tensor_sfb.stride[0][0], ((x, x), x_times_3)),
                tma_tensor_sfb.stride[1],
                tma_tensor_sfb.stride[2],
            )
            tma_tensor_sfb_new_layout = cute.make_layout(new_shape, stride=new_stride)
            tma_tensor_sfb = cute.make_tensor(
                tma_tensor_sfb.iterator, tma_tensor_sfb_new_layout
            )

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

        epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
        tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
            cpasync.CopyBulkTensorTileS2GOp(),
            c_tensor,
            epi_smem_layout,
            self.epi_tile,
        )

        self.tile_sched_params, grid = self._compute_grid(
            c_tensor,
            self.cta_tile_shape_mnk,
            self.cluster_shape_mn,
            max_active_clusters,
        )

        self.buffer_align_bytes = 1024

        @cute.struct
        class SharedStorage:
            ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
            acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
            tmem_dealloc_mbar_ptr: cutlass.Int64
            tmem_holding_buf: cutlass.Int32

            sC: cute.struct.Align[
                cute.struct.MemRange[
                    self.c_dtype,
                    cute.cosize(self.c_smem_layout_staged.outer),
                ],
                self.buffer_align_bytes,
            ]

            sA: cute.struct.Align[
                cute.struct.MemRange[
                    self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]

            sB: cute.struct.Align[
                cute.struct.MemRange[
                    self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
                ],
                self.buffer_align_bytes,
            ]

            sSFA: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]

            sSFB: cute.struct.Align[
                cute.struct.MemRange[
                    self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
                ],
                self.buffer_align_bytes,
            ]

        self.shared_storage = SharedStorage

        grid = (1, 1, 128) # TODO: Blackwell has 148 SM, but 128 is better for some reason

        self.kernel(
            tiled_mma,
            tiled_mma_sfb,
            tma_atom_a,
            tma_tensor_a,
            tma_atom_b,
            tma_tensor_b,
            tma_atom_sfa,
            tma_tensor_sfa,
            tma_atom_sfb,
            tma_tensor_sfb,
            tma_atom_c,
            tma_tensor_c,
            self.cluster_layout_vmnk,
            self.cluster_layout_sfb_vmnk,
            self.a_smem_layout_staged,
            self.b_smem_layout_staged,
            self.sfa_smem_layout_staged,
            self.sfb_smem_layout_staged,
            self.c_smem_layout_staged,
            self.epi_tile,
            self.tile_sched_params,
            epilogue_op,
        ).launch(
            grid=grid,
            block=[self.threads_per_cta, 1, 1],
            cluster=(*self.cluster_shape_mn, 1),
            stream=stream,
        )
        return

    @cute.kernel
    def kernel(
        self,
        tiled_mma: cute.TiledMma,
        tiled_mma_sfb: cute.TiledMma,
        tma_atom_a: cute.CopyAtom,
        mA_mkl: cute.Tensor,
        tma_atom_b: cute.CopyAtom,
        mB_nkl: cute.Tensor,
        tma_atom_sfa: cute.CopyAtom,
        mSFA_mkl: cute.Tensor,
        tma_atom_sfb: cute.CopyAtom,
        mSFB_nkl: cute.Tensor,
        tma_atom_c: cute.CopyAtom,
        mC_mnl: cute.Tensor,
        cluster_layout_vmnk: cute.Layout,
        cluster_layout_sfb_vmnk: cute.Layout,
        a_smem_layout_staged: cute.ComposedLayout,
        b_smem_layout_staged: cute.ComposedLayout,
        sfa_smem_layout_staged: cute.Layout,
        sfb_smem_layout_staged: cute.Layout,
        c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
        epi_tile: cute.Tile,
        tile_sched_params: PersistentTileSchedulerParams,
        epilogue_op: cutlass.Constexpr,
    ):
        warp_idx = cute.arch.warp_idx()
        warp_idx = cute.arch.make_warp_uniform(warp_idx)

        if warp_idx == self.tma_warp_id:
            cpasync.prefetch_descriptor(tma_atom_a)
            cpasync.prefetch_descriptor(tma_atom_b)
            cpasync.prefetch_descriptor(tma_atom_sfa)
            cpasync.prefetch_descriptor(tma_atom_sfb)
            cpasync.prefetch_descriptor(tma_atom_c)

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

        bidx, bidy, bidz = cute.arch.block_idx()
        mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
        is_leader_cta = mma_tile_coord_v == 0
        cta_rank_in_cluster = cute.arch.make_warp_uniform(
            cute.arch.block_idx_in_cluster()
        )
        block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
            cta_rank_in_cluster
        )
        block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
            cta_rank_in_cluster
        )

        tidx, _, _ = cute.arch.thread_idx()

        smem = utils.SmemAllocator()
        storage = smem.allocate(self.shared_storage)

        ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
        ab_pipeline_consumer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread, num_tma_producer
        )
        ab_pipeline = pipeline.PipelineTmaUmma.create(
            barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
            num_stages=self.num_ab_stage,
            producer_group=ab_pipeline_producer_group,
            consumer_group=ab_pipeline_consumer_group,
            tx_count=self.num_tma_load_bytes,
            cta_layout_vmnk=cluster_layout_vmnk,
        )

        acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
        num_acc_consumer_threads = len(self.epilog_warp_id) * (
            2 if use_2cta_instrs else 1
        )
        acc_pipeline_consumer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread, num_acc_consumer_threads
        )
        acc_pipeline = pipeline.PipelineUmmaAsync.create(
            barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
            num_stages=self.num_acc_stage,
            producer_group=acc_pipeline_producer_group,
            consumer_group=acc_pipeline_consumer_group,
            cta_layout_vmnk=cluster_layout_vmnk,
        )

        tmem = utils.TmemAllocator(
            storage.tmem_holding_buf,
            barrier_for_retrieve=self.tmem_alloc_barrier,
            allocator_warp_id=self.epilog_warp_id[0],
            is_two_cta=use_2cta_instrs,
            two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
        )

        if cute.size(self.cluster_shape_mn) > 1:
            cute.arch.cluster_arrive_relaxed()

        sC = storage.sC.get_tensor(
            c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
        )

        sA = storage.sA.get_tensor(
            a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
        )

        sB = storage.sB.get_tensor(
            b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
        )

        sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)

        sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)

        a_full_mcast_mask = None
        b_full_mcast_mask = None
        sfa_full_mcast_mask = None
        sfb_full_mcast_mask = None
        if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
            a_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )
            b_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
            )
            sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
            )
            sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(
                cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1
            )

        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )

        gB_nkl = cute.local_tile(
            mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
        )

        gSFA_mkl = cute.local_tile(
            mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
        )

        gSFB_nkl = cute.local_tile(
            mSFB_nkl,
            cute.slice_(self.mma_tiler_sfb, (0, None, None)),
            (None, None, None),
        )

        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
        )
        k_tile_cnt = cute.size(gA_mkl, mode=[3])

        thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
        thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)

        tCgA = thr_mma.partition_A(gA_mkl)

        tCgB = thr_mma.partition_B(gB_nkl)

        tCgSFA = thr_mma.partition_A(gSFA_mkl)

        tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)

        tCgC = thr_mma.partition_C(gC_mnl)

        a_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
        )

        tAsA, tAgA = cpasync.tma_partition(
            tma_atom_a,
            block_in_cluster_coord_vmnk[2],
            a_cta_layout,
            cute.group_modes(sA, 0, 3),
            cute.group_modes(tCgA, 0, 3),
        )

        b_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
        )

        tBsB, tBgB = cpasync.tma_partition(
            tma_atom_b,
            block_in_cluster_coord_vmnk[1],
            b_cta_layout,
            cute.group_modes(sB, 0, 3),
            cute.group_modes(tCgB, 0, 3),
        )

        sfa_cta_layout = a_cta_layout

        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)

        sfb_cta_layout = cute.make_layout(
            cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
        )

        tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition(
            tma_atom_sfb,
            block_in_cluster_coord_sfb_vmnk[1],
            sfb_cta_layout,
            cute.group_modes(sSFB, 0, 3),
            cute.group_modes(tCgSFB, 0, 3),
        )
        tBsSFB = cute.filter_zeros(tBsSFB)
        tBgSFB = cute.filter_zeros(tBgSFB)

        tCrA = tiled_mma.make_fragment_A(sA)

        tCrB = tiled_mma.make_fragment_B(sB)

        acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])

        tCtAcc_fake = tiled_mma.make_fragment_C(
            cute.append(acc_shape, self.num_acc_stage)
        )

        if cute.size(self.cluster_shape_mn) > 1:
            cute.arch.cluster_wait()
        else:
            self.cta_sync_barrier.arrive_and_wait()

        if warp_idx == self.tma_warp_id:
            tile_sched = StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )

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

            work_batch_size = 1
            done = False
            while not done:
                work_batch = tile_sched.acquire_work_batch(
                    batch_size=work_batch_size
                )
                tiles_processed = 0
                for work_tile in work_batch:
                    if done or not work_tile.is_valid_tile:
                        done = True
                    else:
                        tiles_processed += 1

                        cur_tile_coord = work_tile.tile_idx
                        mma_tile_coord_mnl = (
                            cur_tile_coord[0]
                            // cute.size(tiled_mma.thr_id.shape),
                            cur_tile_coord[1],
                            cur_tile_coord[2],
                        )

                        tAgA_slice = tAgA[
                            (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                        ]

                        tBgB_slice = tBgB[
                            (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
                        ]

                        tAgSFA_slice = tAgSFA[
                            (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
                        ]

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

                        tBgSFB_slice = tBgSFB[
                            (None, slice_n, None, mma_tile_coord_mnl[2])
                        ]

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

                        for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
                            ab_pipeline.producer_acquire(
                                ab_producer_state, peek_ab_empty_status
                            )

                            cute.copy(
                                tma_atom_a,
                                tAgA_slice[(None, ab_producer_state.count)],
                                tAsA[(None, ab_producer_state.index)],
                                tma_bar_ptr=ab_pipeline.producer_get_barrier(
                                    ab_producer_state
                                ),
                                mcast_mask=a_full_mcast_mask,
                            )
                            cute.copy(
                                tma_atom_b,
                                tBgB_slice[(None, ab_producer_state.count)],
                                tBsB[(None, ab_producer_state.index)],
                                tma_bar_ptr=ab_pipeline.producer_get_barrier(
                                    ab_producer_state
                                ),
                                mcast_mask=b_full_mcast_mask,
                            )
                            cute.copy(
                                tma_atom_sfa,
                                tAgSFA_slice[(None, ab_producer_state.count)],
                                tAsSFA[(None, ab_producer_state.index)],
                                tma_bar_ptr=ab_pipeline.producer_get_barrier(
                                    ab_producer_state
                                ),
                                mcast_mask=sfa_full_mcast_mask,
                            )
                            cute.copy(
                                tma_atom_sfb,
                                tBgSFB_slice[(None, ab_producer_state.count)],
                                tBsSFB[(None, ab_producer_state.index)],
                                tma_bar_ptr=ab_pipeline.producer_get_barrier(
                                    ab_producer_state
                                ),
                                mcast_mask=sfb_full_mcast_mask,
                            )

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

                if tiles_processed:
                    tile_sched.advance_to_next_work(advance_count=tiles_processed)

            ab_pipeline.producer_tail(ab_producer_state)

        if warp_idx == self.mma_warp_id:
            tmem.wait_for_alloc()

            acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)

            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

            sfa_tmem_ptr = cute.recast_ptr(
                acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base),
                dtype=self.sf_dtype,
            )

            tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
                tiled_mma,
                self.mma_tiler,
                self.sf_vec_size,
                cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
            )
            tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)

            sfb_tmem_ptr = cute.recast_ptr(
                acc_tmem_ptr
                + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
                + tcgen05.find_tmem_tensor_col_offset(tCtSFA),
                dtype=self.sf_dtype,
            )

            tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
                tiled_mma,
                self.mma_tiler,
                self.sf_vec_size,
                cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
            )
            tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)

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

            tile_sched = StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )

            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
            )

            work_batch_size = 1
            done = False
            while not done:
                work_batch = tile_sched.acquire_work_batch(
                    batch_size=work_batch_size
                )
                tiles_processed = 0
                for work_tile in work_batch:
                    if done or not work_tile.is_valid_tile:
                        done = True
                    else:
                        tiles_processed += 1

                        cur_tile_coord = work_tile.tile_idx
                        mma_tile_coord_mnl = (
                            cur_tile_coord[0]
                            // cute.size(tiled_mma.thr_id.shape),
                            cur_tile_coord[1],
                            cur_tile_coord[2],
                        )

                        tCtAcc = tCtAcc_base[
                            (None, None, None, acc_producer_state.index)
                        ]

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

                        if is_leader_cta:
                            acc_pipeline.producer_acquire(acc_producer_state)

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

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

                        for k_tile in range(k_tile_cnt):
                            if is_leader_cta:
                                ab_pipeline.consumer_wait(
                                    ab_consumer_state, peek_ab_full_status
                                )

                                s2t_stage_coord = (
                                    None,
                                    None,
                                    None,
                                    None,
                                    ab_consumer_state.index,
                                )
                                tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[
                                    s2t_stage_coord
                                ]
                                tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[
                                    s2t_stage_coord
                                ]
                                cute.copy(
                                    tiled_copy_s2t_sfa,
                                    tCsSFA_compact_s2t_staged,
                                    tCtSFA_compact_s2t,
                                )
                                cute.copy(
                                    tiled_copy_s2t_sfb,
                                    tCsSFB_compact_s2t_staged,
                                    tCtSFB_compact_s2t,
                                )

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

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

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

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

                                ab_pipeline.consumer_release(ab_consumer_state)

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

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

                if tiles_processed:
                    tile_sched.advance_to_next_work(advance_count=tiles_processed)

            acc_pipeline.producer_tail(acc_producer_state)

        if warp_idx < self.mma_warp_id:
            tmem.allocate(self.num_tmem_alloc_cols)

            tmem.wait_for_alloc()

            acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)

            tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

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

            tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
            tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
                tiled_copy_t2r, tTR_rC, epi_tidx, sC
            )
            (
                tma_atom_c,
                bSG_sC,
                bSG_gC_partitioned,
            ) = self.epilog_gmem_copy_and_partition(
                epi_tidx, tma_atom_c, tCgC, epi_tile, sC
            )

            tile_sched = StaticPersistentTileScheduler.create(
                tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
            )

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

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

            work_batch_size = 1
            done = False
            while not done:
                work_batch = tile_sched.acquire_work_batch(
                    batch_size=work_batch_size
                )
                tiles_processed = 0
                for work_tile in work_batch:
                    if done or not work_tile.is_valid_tile:
                        done = True
                    else:
                        tile_offset_in_batch = tiles_processed
                        tiles_processed += 1

                        cur_tile_coord = work_tile.tile_idx
                        mma_tile_coord_mnl = (
                            cur_tile_coord[0]
                            // cute.size(tiled_mma.thr_id.shape),
                            cur_tile_coord[1],
                            cur_tile_coord[2],
                        )

                        bSG_gC = bSG_gC_partitioned[
                            (
                                None,
                                None,
                                None,
                                *mma_tile_coord_mnl,
                            )
                        ]

                        tTR_tAcc = tTR_tAcc_base[
                            (None, None, None, None, None, acc_consumer_state.index)
                        ]

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

                        subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
                        num_prev_subtiles = (
                            (tile_sched.num_tiles_executed + Int32(tile_offset_in_batch))
                            * subtile_cnt
                        )
                        for subtile_idx in cutlass.range(subtile_cnt):
                            tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
                            cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

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

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

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

                            if warp_idx == self.epilog_warp_id[0]:
                                cute.copy(
                                    tma_atom_c,
                                    bSG_sC[(None, c_buffer)],
                                    bSG_gC[(None, subtile_idx)],
                                )

                                c_pipeline.producer_commit()
                                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()

                if tiles_processed:
                    tile_sched.advance_to_next_work(advance_count=tiles_processed)

            tmem.relinquish_alloc_permit()
            self.epilog_sync_barrier.arrive_and_wait()
            tmem.free(acc_tmem_ptr)

            c_pipeline.producer_tail()

    def mainloop_s2t_copy_and_partition(
        self,
        sSF: cute.Tensor,
        tSF: cute.Tensor,
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        tCsSF_compact = cute.filter_zeros(sSF)

        tCtSF_compact = cute.filter_zeros(tSF)

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

        tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)

        tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
            tiled_copy_s2t, tCsSF_compact_s2t_
        )

        tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)

        return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t

    def epilog_tmem_copy_and_partition(
        self,
        tidx: cutlass.Int32,
        tAcc: cute.Tensor,
        gC_mnl: cute.Tensor,
        epi_tile: cute.Tile,
        use_2cta_instrs: Union[cutlass.Boolean, bool],
    ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
        copy_atom_t2r = sm100_utils.get_tmem_load_op(
            self.cta_tile_shape_mnk,
            self.c_layout,
            self.c_dtype,
            self.acc_dtype,
            epi_tile,
            use_2cta_instrs,
        )

        tAcc_epi = cute.flat_divide(
            tAcc[((None, None), 0, 0, None)],
            epi_tile,
        )

        tiled_copy_t2r = tcgen05.make_tmem_copy(
            copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
        )

        thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)

        tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)

        gC_mnl_epi = cute.flat_divide(
            gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
        )

        tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)

        tTR_rAcc = cute.make_rmem_tensor(
            tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
        )
        return tiled_copy_t2r, tTR_tAcc, tTR_rAcc

    def epilog_smem_copy_and_partition(
        self,
        tiled_copy_t2r: 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.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
        )
        tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)

        thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
        tRS_sC = thr_copy_r2s.partition_D(sC)

        tRS_rC = tiled_copy_r2s.retile(tTR_rC)
        return tiled_copy_r2s, tRS_rC, tRS_sC

    def epilog_gmem_copy_and_partition(
        self,
        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]:
        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)

        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

    @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,
        occupancy: int,
    ) -> Tuple[int, int, int]:
        num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2

        num_c_stage = 2

        a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
            tiled_mma,
            mma_tiler_mnk,
            a_dtype,
            1,
        )
        b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
            tiled_mma,
            mma_tiler_mnk,
            b_dtype,
            1,
        )
        sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,
        )
        sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
            tiled_mma,
            mma_tiler_mnk,
            sf_vec_size,
            1,
        )

        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

        num_ab_stage = (
            smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
        ) // ab_bytes_per_stage

        num_c_stage += (
            smem_capacity
            - occupancy * ab_bytes_per_stage * num_ab_stage
            - occupancy * (mbar_helpers_bytes + c_bytes)
        ) // (occupancy * c_bytes_per_stage)

        return num_acc_stage, num_ab_stage, num_c_stage

    @staticmethod
    def _compute_grid(
        c: cute.Tensor,
        cta_tile_shape_mnk: Tuple[int, int, int],
        cluster_shape_mn: Tuple[int, int],
        max_active_clusters: cutlass.Constexpr,
    ) -> Tuple[PersistentTileSchedulerParams, Tuple[int, int, int]]:
        c_shape = cute.slice_(cta_tile_shape_mnk, (None, None, 0))
        gc = cute.zipped_divide(c, tiler=c_shape)
        num_ctas_mnl = gc[(0, (None, None, None))].shape
        cluster_shape_mnl = (*cluster_shape_mn, 1)

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

        return tile_sched_params, grid


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

_compiled_kernel_cache = None


def compile_kernel(current_stream):
    global _compiled_kernel_cache

    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache

    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)

    cluster_shape_mn = (1, 1)
    kernel = OptimizedKernel(
        sf_vec_size=16, mma_tiler_mn=(128, 8), cluster_shape_mn=cluster_shape_mn
    )

    hardware_info = cutlass.utils.HardwareInfo()
    max_active_clusters = hardware_info.get_max_active_clusters(
        cluster_shape_mn[0] * cluster_shape_mn[1]
    )
    # TODO: Gives better timing when hardcoded weirdly.
    max_active_clusters = 128

    _compiled_kernel_cache = cute.compile(
        kernel,
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
        (0, 0, 0, 0),
        max_active_clusters,
        current_stream,
        # options="--generate-line-info",
    )

    return _compiled_kernel_cache


def custom_kernel(data: input_t) -> output_t:
    a, b, _, _, sfa_permuted, sfb_permuted, c = data

    current_stream = cutlass_torch.default_stream()
    compiled_func = compile_kernel(current_stream)

    m, k, l = a.shape
    k = k * 2
    n = 1

    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
    )

    compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l), current_stream)

    return c
scrolls · 1878 lines total

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

Changes from previous submission

Against this author's previous submission submission 83205.

⋯ 294 unchanged lines
)
self._problem_cluster_shape = None
- self._mn_cluster_plane = None
if self.params.swizzle_size == 1:
layout_shape = self.params.problem_layout_ncluster_mnl.shape
m_dim, n_dim, l_dim = layout_shape
⋯ 2 unchanged lines
Int32(cute.size(n_dim)),
Int32(cute.size(l_dim)),
)
- self._mn_cluster_plane = (
- self._problem_cluster_shape[0] * self._problem_cluster_shape[1]
- )
def __extract_mlir_values__(self) -> list[ir.Value]:
values = extract_mlir_values(self.num_persistent_clusters)
⋯ 160 unchanged lines
)
self._num_tiles_executed += advance
+ @dsl_user_op
+ def acquire_work_batch(self, *, batch_size: int = 1, loc=None, ip=None):
+ if batch_size < 1:
+ raise ValueError("batch_size must be >= 1")
+
+ tiles: list[WorkTileInfo] = []
+ base_linear_idx = self._current_work_linear_idx
+ stride = Int32(self.num_persistent_clusters)
+
+ for batch_idx in range(batch_size):
+ tile_linear_idx = base_linear_idx + Int32(batch_idx) * stride
+ tiles.append(
+ self._get_current_work_for_linear_idx(
+ tile_linear_idx, loc=loc, ip=ip
+ )
+ )
+
+ return tuple(tiles)
+
@property
def num_tiles_executed(self) -> Int32:
return self._num_tiles_executed
⋯ 657 unchanged lines
tile_sched = StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
- work_tile = tile_sched.initial_work_tile_info()
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
- while work_tile.is_valid_tile:
- cur_tile_coord = work_tile.tile_idx
- mma_tile_coord_mnl = (
- cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
- cur_tile_coord[1],
- cur_tile_coord[2],
+ work_batch_size = 1
+ done = False
+ while not done:
+ work_batch = tile_sched.acquire_work_batch(
+ batch_size=work_batch_size
)
+ tiles_processed = 0
+ for work_tile in work_batch:
+ if done or not work_tile.is_valid_tile:
+ done = True
+ else:
+ tiles_processed += 1
- tAgA_slice = tAgA[
- (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
- ]
+ cur_tile_coord = work_tile.tile_idx
+ mma_tile_coord_mnl = (
+ cur_tile_coord[0]
+ // cute.size(tiled_mma.thr_id.shape),
+ cur_tile_coord[1],
+ cur_tile_coord[2],
+ )
- tBgB_slice = tBgB[
- (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
- ]
+ tAgA_slice = tAgA[
+ (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
+ ]
- tAgSFA_slice = tAgSFA[
- (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
- ]
+ tBgB_slice = tBgB[
+ (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
+ ]
- slice_n = mma_tile_coord_mnl[1]
- if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
- slice_n = mma_tile_coord_mnl[1] // 2
+ tAgSFA_slice = tAgSFA[
+ (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
+ ]
- tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]
+ slice_n = mma_tile_coord_mnl[1]
+ if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
+ slice_n = mma_tile_coord_mnl[1] // 2
- ab_producer_state.reset_count()
- peek_ab_empty_status = cutlass.Boolean(1)
- if ab_producer_state.count < k_tile_cnt:
- peek_ab_empty_status = ab_pipeline.producer_try_acquire(
- ab_producer_state
- )
+ tBgSFB_slice = tBgSFB[
+ (None, slice_n, None, mma_tile_coord_mnl[2])
+ ]
- for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
- ab_pipeline.producer_acquire(
- ab_producer_state, peek_ab_empty_status
- )
+ ab_producer_state.reset_count()
+ peek_ab_empty_status = cutlass.Boolean(1)
+ if ab_producer_state.count < k_tile_cnt:
+ peek_ab_empty_status = ab_pipeline.producer_try_acquire(
+ ab_producer_state
+ )
- cute.copy(
- tma_atom_a,
- tAgA_slice[(None, ab_producer_state.count)],
- tAsA[(None, ab_producer_state.index)],
- tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
- mcast_mask=a_full_mcast_mask,
- )
- cute.copy(
- tma_atom_b,
- tBgB_slice[(None, ab_producer_state.count)],
- tBsB[(None, ab_producer_state.index)],
- tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
- mcast_mask=b_full_mcast_mask,
- )
- cute.copy(
- tma_atom_sfa,
- tAgSFA_slice[(None, ab_producer_state.count)],
- tAsSFA[(None, ab_producer_state.index)],
- tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
- mcast_mask=sfa_full_mcast_mask,
- )
- cute.copy(
- tma_atom_sfb,
- tBgSFB_slice[(None, ab_producer_state.count)],
- tBsSFB[(None, ab_producer_state.index)],
- tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
- mcast_mask=sfb_full_mcast_mask,
- )
+ for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
+ ab_pipeline.producer_acquire(
+ ab_producer_state, peek_ab_empty_status
+ )
- ab_producer_state.advance()
- peek_ab_empty_status = cutlass.Boolean(1)
- if ab_producer_state.count < k_tile_cnt:
- peek_ab_empty_status = ab_pipeline.producer_try_acquire(
- ab_producer_state
- )
+ cute.copy(
+ tma_atom_a,
+ tAgA_slice[(None, ab_producer_state.count)],
+ tAsA[(None, ab_producer_state.index)],
+ tma_bar_ptr=ab_pipeline.producer_get_barrier(
+ ab_producer_state
+ ),
+ mcast_mask=a_full_mcast_mask,
+ )
+ cute.copy(
+ tma_atom_b,
+ tBgB_slice[(None, ab_producer_state.count)],
+ tBsB[(None, ab_producer_state.index)],
+ tma_bar_ptr=ab_pipeline.producer_get_barrier(
+ ab_producer_state
+ ),
+ mcast_mask=b_full_mcast_mask,
+ )
+ cute.copy(
+ tma_atom_sfa,
+ tAgSFA_slice[(None, ab_producer_state.count)],
+ tAsSFA[(None, ab_producer_state.index)],
+ tma_bar_ptr=ab_pipeline.producer_get_barrier(
+ ab_producer_state
+ ),
+ mcast_mask=sfa_full_mcast_mask,
+ )
+ cute.copy(
+ tma_atom_sfb,
+ tBgSFB_slice[(None, ab_producer_state.count)],
+ tBsSFB[(None, ab_producer_state.index)],
+ tma_bar_ptr=ab_pipeline.producer_get_barrier(
+ ab_producer_state
+ ),
+ mcast_mask=sfb_full_mcast_mask,
+ )
- tile_sched.advance_to_next_work()
- work_tile = tile_sched.get_current_work()
+ ab_producer_state.advance()
+ peek_ab_empty_status = cutlass.Boolean(1)
+ if ab_producer_state.count < k_tile_cnt:
+ peek_ab_empty_status = (
+ ab_pipeline.producer_try_acquire(
+ ab_producer_state
+ )
+ )
+ if tiles_processed:
+ tile_sched.advance_to_next_work(advance_count=tiles_processed)
+
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
⋯ 45 unchanged lines
tile_sched = StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
- work_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
⋯ 2 unchanged lines
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
- while work_tile.is_valid_tile:
- cur_tile_coord = work_tile.tile_idx
- mma_tile_coord_mnl = (
- cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
- cur_tile_coord[1],
- cur_tile_coord[2],
+ work_batch_size = 1
+ done = False
+ while not done:
+ work_batch = tile_sched.acquire_work_batch(
+ batch_size=work_batch_size
)
+ tiles_processed = 0
+ for work_tile in work_batch:
+ if done or not work_tile.is_valid_tile:
+ done = True
+ else:
+ tiles_processed += 1
- tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
+ cur_tile_coord = work_tile.tile_idx
+ mma_tile_coord_mnl = (
+ cur_tile_coord[0]
+ // cute.size(tiled_mma.thr_id.shape),
+ cur_tile_coord[1],
+ cur_tile_coord[2],
+ )
- ab_consumer_state.reset_count()
- peek_ab_full_status = cutlass.Boolean(1)
- if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
- peek_ab_full_status = ab_pipeline.consumer_try_wait(
- ab_consumer_state
- )
+ tCtAcc = tCtAcc_base[
+ (None, None, None, acc_producer_state.index)
+ ]
- if is_leader_cta:
- acc_pipeline.producer_acquire(acc_producer_state)
+ ab_consumer_state.reset_count()
+ peek_ab_full_status = cutlass.Boolean(1)
+ if (
+ ab_consumer_state.count < k_tile_cnt and is_leader_cta
+ ):
+ peek_ab_full_status = ab_pipeline.consumer_try_wait(
+ ab_consumer_state
+ )
- tCtSFB_mma = tCtSFB
- if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
- offset = (
- cutlass.Int32(2)
- if mma_tile_coord_mnl[1] % 2 == 1
- else cutlass.Int32(0)
- )
- shifted_ptr = cute.recast_ptr(
- acc_tmem_ptr
- + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
- + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
- + offset,
- dtype=self.sf_dtype,
- )
- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
- elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
- offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
- shifted_ptr = cute.recast_ptr(
- acc_tmem_ptr
- + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
- + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
- + offset,
- dtype=self.sf_dtype,
- )
- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
+ if is_leader_cta:
+ acc_pipeline.producer_acquire(acc_producer_state)
- tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
+ tCtSFB_mma = tCtSFB
+ if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
+ offset = (
+ cutlass.Int32(2)
+ if mma_tile_coord_mnl[1] % 2 == 1
+ else cutlass.Int32(0)
+ )
+ shifted_ptr = cute.recast_ptr(
+ acc_tmem_ptr
+ + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
+ + offset,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
+ elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
+ offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
+ shifted_ptr = cute.recast_ptr(
+ acc_tmem_ptr
+ + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
+ + offset,
+ dtype=self.sf_dtype,
+ )
+ tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
- for k_tile in range(k_tile_cnt):
- if is_leader_cta:
- ab_pipeline.consumer_wait(
- ab_consumer_state, peek_ab_full_status
- )
+ tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
- s2t_stage_coord = (
- None,
- None,
- None,
- None,
- ab_consumer_state.index,
- )
- tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
- tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
- cute.copy(
- tiled_copy_s2t_sfa,
- tCsSFA_compact_s2t_staged,
- tCtSFA_compact_s2t,
- )
- cute.copy(
- tiled_copy_s2t_sfb,
- tCsSFB_compact_s2t_staged,
- tCtSFB_compact_s2t,
- )
+ for k_tile in range(k_tile_cnt):
+ if is_leader_cta:
+ ab_pipeline.consumer_wait(
+ ab_consumer_state, peek_ab_full_status
+ )
- num_kblocks = cute.size(tCrA, mode=[2])
- for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
- kblock_coord = (
- None,
- None,
- kblock_idx,
- ab_consumer_state.index,
- )
+ s2t_stage_coord = (
+ None,
+ None,
+ None,
+ None,
+ ab_consumer_state.index,
+ )
+ tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[
+ s2t_stage_coord
+ ]
+ tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[
+ s2t_stage_coord
+ ]
+ cute.copy(
+ tiled_copy_s2t_sfa,
+ tCsSFA_compact_s2t_staged,
+ tCtSFA_compact_s2t,
+ )
+ cute.copy(
+ tiled_copy_s2t_sfb,
+ tCsSFB_compact_s2t_staged,
+ tCtSFB_compact_s2t,
+ )
- sf_kblock_coord = (None, None, kblock_idx)
- tiled_mma.set(
- tcgen05.Field.SFA,
- tCtSFA[sf_kblock_coord].iterator,
- )
- tiled_mma.set(
- tcgen05.Field.SFB,
- tCtSFB_mma[sf_kblock_coord].iterator,
- )
+ num_kblocks = cute.size(tCrA, mode=[2])
+ for kblock_idx in cutlass.range(
+ num_kblocks, unroll_full=True
+ ):
+ kblock_coord = (
+ None,
+ None,
+ kblock_idx,
+ ab_consumer_state.index,
+ )
- cute.gemm(
- tiled_mma,
- tCtAcc,
- tCrA[kblock_coord],
- tCrB[kblock_coord],
- tCtAcc,
- )
+ sf_kblock_coord = (None, None, kblock_idx)
+ tiled_mma.set(
+ tcgen05.Field.SFA,
+ tCtSFA[sf_kblock_coord].iterator,
+ )
+ tiled_mma.set(
+ tcgen05.Field.SFB,
+ tCtSFB_mma[sf_kblock_coord].iterator,
+ )
- tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
+ cute.gemm(
+ tiled_mma,
+ tCtAcc,
+ tCrA[kblock_coord],
+ tCrB[kblock_coord],
+ tCtAcc,
+ )
- ab_pipeline.consumer_release(ab_consumer_state)
+ tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
- ab_consumer_state.advance()
- peek_ab_full_status = cutlass.Boolean(1)
- if ab_consumer_state.count < k_tile_cnt:
- if is_leader_cta:
- peek_ab_full_status = ab_pipeline.consumer_try_wait(
- ab_consumer_state
- )
+ ab_pipeline.consumer_release(ab_consumer_state)
- if is_leader_cta:
- acc_pipeline.producer_commit(acc_producer_state)
- acc_producer_state.advance()
+ ab_consumer_state.advance()
+ peek_ab_full_status = cutlass.Boolean(1)
+ if ab_consumer_state.count < k_tile_cnt:
+ if is_leader_cta:
+ peek_ab_full_status = (
+ ab_pipeline.consumer_try_wait(
+ ab_consumer_state
+ )
+ )
- tile_sched.advance_to_next_work()
- work_tile = tile_sched.get_current_work()
+ if is_leader_cta:
+ acc_pipeline.producer_commit(acc_producer_state)
+ acc_producer_state.advance()
+ if tiles_processed:
+ tile_sched.advance_to_next_work(advance_count=tiles_processed)
+
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
⋯ 29 unchanged lines
tile_sched = StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
- work_tile = tile_sched.initial_work_tile_info()
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
⋯ 8 unchanged lines
producer_group=c_producer_group,
)
- while work_tile.is_valid_tile:
- cur_tile_coord = work_tile.tile_idx
- mma_tile_coord_mnl = (
- cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
- cur_tile_coord[1],
- cur_tile_coord[2],
+ work_batch_size = 1
+ done = False
+ while not done:
+ work_batch = tile_sched.acquire_work_batch(
+ batch_size=work_batch_size
)
+ tiles_processed = 0
+ for work_tile in work_batch:
+ if done or not work_tile.is_valid_tile:
+ done = True
+ else:
+ tile_offset_in_batch = tiles_processed
+ tiles_processed += 1
- bSG_gC = bSG_gC_partitioned[
- (
- None,
- None,
- None,
- *mma_tile_coord_mnl,
- )
- ]
+ cur_tile_coord = work_tile.tile_idx
+ mma_tile_coord_mnl = (
+ cur_tile_coord[0]
+ // cute.size(tiled_mma.thr_id.shape),
+ cur_tile_coord[1],
+ cur_tile_coord[2],
+ )
- tTR_tAcc = tTR_tAcc_base[
- (None, None, None, None, None, acc_consumer_state.index)
- ]
+ bSG_gC = bSG_gC_partitioned[
+ (
+ None,
+ None,
+ None,
+ *mma_tile_coord_mnl,
+ )
+ ]
- acc_pipeline.consumer_wait(acc_consumer_state)
+ tTR_tAcc = tTR_tAcc_base[
+ (None, None, None, None, None, acc_consumer_state.index)
+ ]
- tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
- bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
+ acc_pipeline.consumer_wait(acc_consumer_state)
- 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):
- tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
- cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
+ tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
+ bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
- acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
- acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
- tRS_rC.store(acc_vec)
+ subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
+ num_prev_subtiles = (
+ (tile_sched.num_tiles_executed + Int32(tile_offset_in_batch))
+ * subtile_cnt
+ )
+ for subtile_idx in cutlass.range(subtile_cnt):
+ tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
+ cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
- c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
- cute.copy(
- tiled_copy_r2s,
- tRS_rC,
- tRS_sC[(None, None, None, c_buffer)],
- )
+ acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
+ acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
+ tRS_rC.store(acc_vec)
- cute.arch.fence_proxy(
- cute.arch.ProxyKind.async_shared,
- space=cute.arch.SharedSpace.shared_cta,
- )
- self.epilog_sync_barrier.arrive_and_wait()
+ c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
+ cute.copy(
+ tiled_copy_r2s,
+ tRS_rC,
+ tRS_sC[(None, None, None, c_buffer)],
+ )
- if warp_idx == self.epilog_warp_id[0]:
- cute.copy(
- tma_atom_c,
- bSG_sC[(None, c_buffer)],
- bSG_gC[(None, subtile_idx)],
- )
+ cute.arch.fence_proxy(
+ cute.arch.ProxyKind.async_shared,
+ space=cute.arch.SharedSpace.shared_cta,
+ )
+ self.epilog_sync_barrier.arrive_and_wait()
- c_pipeline.producer_commit()
- c_pipeline.producer_acquire()
- self.epilog_sync_barrier.arrive_and_wait()
+ if warp_idx == self.epilog_warp_id[0]:
+ cute.copy(
+ tma_atom_c,
+ bSG_sC[(None, c_buffer)],
+ bSG_gC[(None, subtile_idx)],
+ )
- with cute.arch.elect_one():
- acc_pipeline.consumer_release(acc_consumer_state)
- acc_consumer_state.advance()
+ c_pipeline.producer_commit()
+ c_pipeline.producer_acquire()
+ self.epilog_sync_barrier.arrive_and_wait()
- tile_sched.advance_to_next_work()
- work_tile = tile_sched.get_current_work()
+ with cute.arch.elect_one():
+ acc_pipeline.consumer_release(acc_consumer_state)
+ acc_consumer_state.advance()
+ if tiles_processed:
+ tile_sched.advance_to_next_work(advance_count=tiles_processed)
+
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
scrolls · 639 diff lines total

Best evidence level for this revision: reported

JSON