submission 83205
shellsmile15795 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1799 lines, June 9 Researcher Reciprocity License v1.0.
submission_v13.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-83205?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:d16efd70f7acc7977f1371b79b930bb29bd1937bb2707941faa580c6516b9959
license declaredunknown
license concludedunknown
authorsshellsmile15795
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
self.cta_sync_barrier = pipeline.NamedBarrier(persistent-kernel
class PersistentTileSchedulerParams:shared-memory
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")tcgen05
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONEwarp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
submission_v13.py1799 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
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
self._problem_cluster_shape = (
Int32(cute.size(m_dim)),
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)
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
@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()
)
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],
)
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
)
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
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()
)
work_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
while work_tile.is_valid_tile:
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()
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
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()
)
work_tile = tile_sched.initial_work_tile_info()
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,
)
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],
)
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 * 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()
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
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 · 1799 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 83083.
⋯ 21 unchanged linesimport cutlass.utils.blockscaled_layout as blockscaled_utilsfrom 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+ 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+ self._problem_cluster_shape = (+ Int32(cute.size(m_dim)),+ 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)+ 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++ @property+ def num_tiles_executed(self) -> Int32:+ return self._num_tiles_executed++ #+ # NOTE: Optimized Kernel+ #+class OptimizedKernel:def __init__(self,⋯ 381 unchanged linesself.shared_storage = SharedStorage- grid = (1, 1, 128)+ grid = (1, 1, 128) # TODO: Blackwell has 148 SM, but 128 is better for some reasonself.kernel(tiled_mma,⋯ 49 unchanged linessfb_smem_layout_staged: cute.Layout,c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],epi_tile: cute.Tile,- tile_sched_params: utils.PersistentTileSchedulerParams,+ tile_sched_params: PersistentTileSchedulerParams,epilogue_op: cutlass.Constexpr,):warp_idx = cute.arch.warp_idx()⋯ 202 unchanged linesself.cta_sync_barrier.arrive_and_wait()if warp_idx == self.tma_warp_id:- tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched = StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim())work_tile = tile_sched.initial_work_tile_info()⋯ 127 unchanged linestCtSFB_compact_s2t,) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)- tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched = StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim())work_tile = tile_sched.initial_work_tile_info()⋯ 157 unchanged linesepi_tidx, tma_atom_c, tCgC, epi_tile, sC)- tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched = StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim())work_tile = tile_sched.initial_work_tile_info()⋯ 272 unchanged linescta_tile_shape_mnk: Tuple[int, int, int],cluster_shape_mn: Tuple[int, int],max_active_clusters: cutlass.Constexpr,- ) -> Tuple[utils.PersistentTileSchedulerParams, Tuple[int, int, int]]:+ ) -> 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))].shapecluster_shape_mnl = (*cluster_shape_mn, 1)- tile_sched_params = utils.PersistentTileSchedulerParams(+ tile_sched_params = PersistentTileSchedulerParams(num_ctas_mnl, cluster_shape_mnl)- grid = utils.StaticPersistentTileScheduler.get_grid_shape(+ grid = StaticPersistentTileScheduler.get_grid_shape(tile_sched_params, max_active_clusters)
scrolls · 531 diff lines total
Best evidence level for this revision: reported
JSON