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
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-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_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 = Noneif self.params.swizzle_size == 1:layout_shape = self.params.problem_layout_ncluster_mnl.shapem_dim, n_dim, l_dim = layout_shape⋯ 2 unchanged linesInt32(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)+@propertydef num_tiles_executed(self) -> Int32:return self._num_tiles_executed⋯ 657 unchanged linestile_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 linestile_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 linespipeline.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 linestile_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 linesproducer_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