submission 409213
moshe_05859 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1091 lines, June 9 Researcher Reciprocity License v1.0.
cudeepy_nvfp4_group_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-409213?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:bd926c89c7ced28ce3bd276c5972a726c7e4faff3d4f75f5c9d61756e2bc5826
license declaredunknown
license concludedunknown
authorsmoshe_05859
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(self.cta_tile_shape_mnk, self.use_2cta_instrs, self.c_layout, self.c_dtype)mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(self.epilog_warp_id))persistent-kernel
tile_sched_params: utils.PersistentTileSchedulerParams, group_count: cutlass.Constexpr,shared-memory
reserved_smem_bytes = 1024tcgen05
self.cta_group = tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONEtile-m = 128
CLUSTER_TILE_M = 128 * CLUSTER_Mtile-n = 128
CLUSTER_TILE_M, CLUSTER_TILE_N = 128 * CLUSTER_M, MMA_N * CLUSTER_Nwarp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
cudeepy_nvfp4_group_gemm.py1091 lines
# SM100 Grouped BlockScaled GEMM Kernel (NVFP4) - Condensed Version
# Warp-specialized persistent kernel for Blackwell architecture
import argparse, functools
from typing import List, Type, Tuple, Union
from inspect import isclass
import torch
import cuda.bindings.driver as cuda
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
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import from_dlpack
from cutlass.cutlass_dsl import Boolean, extract_mlir_values, new_from_mlir_values
class GroupedWorkTileInfo:
"""Work tile info for grouped GEMM with group search result."""
def __init__(self, tile_idx: cute.Coord, is_valid_tile: Boolean, group_search_result):
self._tile_idx = tile_idx
self._is_valid_tile = Boolean(is_valid_tile)
self.group_search_result = group_search_result
def __extract_mlir_values__(self) -> list:
values = extract_mlir_values(self.tile_idx)
values.extend(extract_mlir_values(self.is_valid_tile))
values.extend(extract_mlir_values(self.group_search_result))
return values
def __new_from_mlir_values__(self, values: list) -> "GroupedWorkTileInfo":
if len(values) != 11:
raise ValueError("Length of mlir values extracted is incorrect.")
return GroupedWorkTileInfo(
new_from_mlir_values(self._tile_idx, values[:3]),
new_from_mlir_values(self._is_valid_tile, [values[3]]),
new_from_mlir_values(self.group_search_result, values[4:11]),
)
@property
def is_valid_tile(self) -> Boolean:
return self._is_valid_tile
@property
def tile_idx(self) -> cute.Coord:
return self._tile_idx
class Sm100GroupedBlockScaledGemmKernel:
"""SM100 grouped blockscaled GEMM kernel with TMA + warp specialization."""
reserved_smem_bytes = 1024
bytes_per_tensormap = 128
num_tensormaps = 5
tensor_memory_management_bytes = 12
def __init__(self, sf_vec_size: int, mma_tiler_mn: Tuple[int, int], cluster_shape_mn: Tuple[int, int], occupancy: int = 1, max_ab_stages: int = 0):
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.tensormap_update_mode = utils.TensorMapUpdateMode.SMEM
self.occupancy = occupancy
self.max_ab_stages = max_ab_stages # 0 = auto, >0 = limit stages
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.epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(self.epilog_warp_id))
self.tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=2, num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)))
self.tensormap_ab_init_barrier = pipeline.NamedBarrier(barrier_id=3, num_threads=64)
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
self.num_tmem_alloc_cols = 512
def _setup_attributes(self):
"""Configure MMA shapes, cluster layouts, multicast masks, stages, and smem layouts."""
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, 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_tile_shape_mnk = tuple(x * y for x, y in zip(self.cta_tile_shape_mnk, (*self.cluster_shape_mn, 1)))
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,
)
# Apply max_ab_stages limit if set
if self.max_ab_stages > 0 and self.num_ab_stage > self.max_ab_stages:
self.num_ab_stage = self.max_ab_stages
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)
mbar_smem_bytes = self._get_mbar_smem_bytes(num_acc_stage=self.num_acc_stage, num_ab_stage=self.num_ab_stage, num_c_stage=self.num_c_stage)
tensormap_smem_bytes = self.bytes_per_tensormap * self.num_tensormaps
if mbar_smem_bytes + tensormap_smem_bytes + self.tensor_memory_management_bytes > self.reserved_smem_bytes:
raise ValueError(f"smem consumption exceeds reserved {self.reserved_smem_bytes}")
@cute.jit
def __call__(self, initial_a: cute.Tensor, initial_b: cute.Tensor, initial_c: cute.Tensor,
initial_sfa: cute.Tensor, initial_sfb: cute.Tensor, group_count: cutlass.Constexpr[int],
problem_shape_mnkl: cute.Tensor, strides_abc: cute.Tensor, tensor_address_abc: cute.Tensor,
tensor_address_sfasfb: cute.Tensor, total_num_clusters: cutlass.Constexpr[int],
tensormap_cute_tensor: cute.Tensor, max_active_clusters: cutlass.Constexpr[int]):
"""Setup and launch grouped GEMM kernel."""
self.a_dtype, self.b_dtype = initial_a.element_type, initial_b.element_type
self.sf_dtype, self.c_dtype = initial_sfa.element_type, initial_c.element_type
self.is_nvfp4_output = self.c_dtype is cutlass.Float4E2M1FN
self.a_major_mode = utils.LayoutEnum.from_tensor(initial_a).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(initial_b).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(initial_c)
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type mismatch: {self.a_dtype} != {self.b_dtype}")
self._setup_attributes()
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_a.shape, self.sf_vec_size)
initial_sfa = cute.make_tensor(initial_sfa.iterator, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_b.shape, self.sf_vec_size)
initial_sfb = cute.make_tensor(initial_sfb.iterator, sfb_layout)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
self.sf_vec_size, self.cta_group, self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype, self.a_major_mode, self.b_major_mode, self.sf_dtype,
self.sf_vec_size, tcgen05.CtaGroup.ONE, self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
# TMA setup for A/B/SFA/SFB/C
a_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(a_op, initial_a, 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, initial_b, 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, initial_sfa, 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, initial_sfb, sfb_smem_layout, self.mma_tiler_sfb, tiled_mma_sfb, self.cluster_layout_sfb_vmnk.shape, internal_type=cutlass.Int16)
self.num_tma_load_bytes = (cute.size_in_bytes(self.a_dtype, a_smem_layout) + cute.size_in_bytes(self.b_dtype, b_smem_layout) +
cute.size_in_bytes(self.sf_dtype, sfa_smem_layout) + cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)) * 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(), initial_c, epi_smem_layout, self.epi_tile)
self.tile_sched_params, grid = self._compute_grid(total_num_clusters, self.cluster_shape_mn, max_active_clusters)
self.buffer_align_bytes = 1024
self.size_tensormap_in_i64 = self.num_tensormaps * self.bytes_per_tensormap // 8
@cute.struct
class SharedStorage:
tensormap_buffer: cute.struct.MemRange[cutlass.Int64, self.size_tensormap_in_i64]
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[cute.struct.MemRange[self.c_dtype, cute.cosize(self.c_smem_layout_staged.outer)], self.buffer_align_bytes]
sA: cute.struct.Align[cute.struct.MemRange[self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)], self.buffer_align_bytes]
sB: cute.struct.Align[cute.struct.MemRange[self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)], self.buffer_align_bytes]
sSFA: cute.struct.Align[cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)], self.buffer_align_bytes]
sSFB: cute.struct.Align[cute.struct.MemRange[self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)], self.buffer_align_bytes]
self.shared_storage = SharedStorage
self.kernel(tiled_mma, tiled_mma_sfb, tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b,
tma_atom_sfa, tma_tensor_sfa, tma_atom_sfb, tma_tensor_sfb, tma_atom_c, tma_tensor_c,
self.cluster_layout_vmnk, self.cluster_layout_sfb_vmnk,
self.a_smem_layout_staged, self.b_smem_layout_staged, self.sfa_smem_layout_staged,
self.sfb_smem_layout_staged, self.c_smem_layout_staged, self.epi_tile,
self.tile_sched_params, group_count, problem_shape_mnkl, strides_abc,
tensor_address_abc, tensor_address_sfasfb, tensormap_cute_tensor,
).launch(grid=grid, block=[self.threads_per_cta, 1, 1], cluster=(*self.cluster_shape_mn, 1),
smem=self.shared_storage.size_in_bytes(), min_blocks_per_mp=1)
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: utils.PersistentTileSchedulerParams, group_count: cutlass.Constexpr,
problem_sizes_mnkl: cute.Tensor, strides_abc: cute.Tensor, ptrs_abc: cute.Tensor,
ptrs_sfasfb: cute.Tensor, tensormaps: cute.Tensor):
"""GPU device kernel - warp-specialized with TMA/MMA/Epilogue warps."""
warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx())
if warp_idx == self.tma_warp_id:
cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_a)
cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_b)
cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_sfa)
cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_sfb)
cute.nvgpu.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()
# Allocate and init shared memory
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
tensormap_smem_ptr = storage.tensormap_buffer.data_ptr()
tensormap_a_smem_ptr = tensormap_smem_ptr
tensormap_b_smem_ptr = tensormap_a_smem_ptr + self.bytes_per_tensormap // 8
tensormap_sfa_smem_ptr = tensormap_b_smem_ptr + self.bytes_per_tensormap // 8
tensormap_sfb_smem_ptr = tensormap_sfa_smem_ptr + self.bytes_per_tensormap // 8
tensormap_c_smem_ptr = tensormap_sfb_smem_ptr + self.bytes_per_tensormap // 8
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
# Init pipelines
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,
)
if use_2cta_instrs:
if warp_idx == self.tma_warp_id:
with cute.arch.elect_one():
cute.arch.mbarrier_init(tmem_dealloc_mbar_ptr, 32)
pipeline_init_arrive(cluster_shape_mn=self.cluster_shape_mn, is_relaxed=True)
# Setup smem tensors
sC = storage.sC.get_tensor(c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner)
sA = storage.sA.get_tensor(a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner)
sB = storage.sB.get_tensor(b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
# Multicast masks
a_full_mcast_mask, b_full_mcast_mask, sfa_full_mcast_mask, sfb_full_mcast_mask = None, None, None, 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)
# Local_tile partition global tensors
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, (0, None, None)), (None, None, None))
gC_mnl = cute.local_tile(mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None))
# Partition for TiledMMA
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgB = thr_mma.partition_B(gB_nkl)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
# TMA partitions
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))
tAsSFA, tAgSFA = cpasync.tma_partition(tma_atom_sfa, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sSFA, 0, 3), cute.group_modes(tCgSFA, 0, 3))
tAsSFA, tAgSFA = cute.filter_zeros(tAsSFA), cute.filter_zeros(tAgSFA)
sfb_cta_layout = cute.make_layout(cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape)
tBsSFB, tBgSFB = cpasync.tma_partition(tma_atom_sfb, block_in_cluster_coord_sfb_vmnk[1], sfb_cta_layout, cute.group_modes(sSFB, 0, 3), cute.group_modes(tCgSFB, 0, 3))
tBsSFB, tBgSFB = cute.filter_zeros(tBsSFB), cute.filter_zeros(tBgSFB)
# Partition for MMA
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))
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
# Tensormap setup
grid_dim = cute.arch.grid_dim()
tensormap_workspace_idx = bidz * grid_dim[1] * grid_dim[0] + bidy * grid_dim[0] + bidx
tensormap_manager = utils.TensorMapManager(utils.TensorMapUpdateMode.SMEM, self.bytes_per_tensormap)
tensormap_a_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 0, None)].iterator)
tensormap_b_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 1, None)].iterator)
tensormap_sfa_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 2, None)].iterator)
tensormap_sfb_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 3, None)].iterator)
tensormap_c_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 4, None)].iterator)
# Tile scheduler setup
tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
group_helper = utils.GroupedGemmTileSchedulerHelper(group_count, tile_sched_params, self.cluster_tile_shape_mnk, utils.create_initial_search_state())
base_work_tile = tile_sched.initial_work_tile_info()
group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
initial_work_tile_info = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
# TMA WARP
if warp_idx == self.tma_warp_id and initial_work_tile_info.is_valid_tile:
work_tile = initial_work_tile_info
tensormap_init_done = cutlass.Boolean(False)
last_group_idx = cutlass.Int32(-1)
ab_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_ab_stage)
while work_tile.is_valid_tile:
grouped_gemm_cta_tile_info = work_tile.group_search_result
cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
cur_group_idx = grouped_gemm_cta_tile_info.group_idx
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
if not is_k_tile_cnt_zero:
is_group_changed = cur_group_idx != last_group_idx
if is_group_changed:
ps = (grouped_gemm_cta_tile_info.problem_shape_m, grouped_gemm_cta_tile_info.problem_shape_n, grouped_gemm_cta_tile_info.problem_shape_k)
real_tensor_a = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.a_dtype, ps, strides_abc, ptrs_abc, 0)
real_tensor_b = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.b_dtype, ps, strides_abc, ptrs_abc, 1)
real_tensor_sfa = self.make_tensor_sfasfb_for_tensormap_update(cur_group_idx, self.sf_dtype, ps, ptrs_sfasfb, 0)
real_tensor_sfb = self.make_tensor_sfasfb_for_tensormap_update(cur_group_idx, self.sf_dtype, ps, ptrs_sfasfb, 1)
if not tensormap_init_done:
self.tensormap_ab_init_barrier.arrive_and_wait()
tensormap_init_done = True
tensormap_manager.update_tensormap(
(real_tensor_a, real_tensor_b, real_tensor_sfa, real_tensor_sfb),
(tma_atom_a, tma_atom_b, tma_atom_sfa, tma_atom_sfb),
(tensormap_a_gmem_ptr, tensormap_b_gmem_ptr, tensormap_sfa_gmem_ptr, tensormap_sfb_gmem_ptr),
self.tma_warp_id, (tensormap_a_smem_ptr, tensormap_b_smem_ptr, tensormap_sfa_smem_ptr, tensormap_sfb_smem_ptr),
)
mma_tile_coord_mnl = (grouped_gemm_cta_tile_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape), grouped_gemm_cta_tile_info.cta_tile_idx_n, 0)
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])]
tBgSFB_slice = tBgSFB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < cur_k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
if is_group_changed:
tensormap_manager.fence_tensormap_update(tensormap_a_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_b_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_sfa_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_sfb_gmem_ptr)
for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
bar_ptr = ab_pipeline.producer_get_barrier(ab_producer_state)
cute.copy(tma_atom_a, tAgA_slice[(None, ab_producer_state.count)], tAsA[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=a_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_a_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_b, tBgB_slice[(None, ab_producer_state.count)], tBsB[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=b_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_b_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_sfa, tAgSFA_slice[(None, ab_producer_state.count)], tAsSFA[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=sfa_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfa_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_sfb, tBgSFB_slice[(None, ab_producer_state.count)], tBsSFB[(None, ab_producer_state.index)], tma_bar_ptr=bar_ptr, mcast_mask=sfb_full_mcast_mask, tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfb_gmem_ptr, cute.AddressSpace.generic))
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < cur_k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
else:
if not tensormap_init_done:
self.tensormap_ab_init_barrier.arrive_and_wait()
tensormap_init_done = True
tile_sched.advance_to_next_work()
base_work_tile = tile_sched.get_current_work()
group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
last_group_idx = cur_group_idx
ab_pipeline.producer_tail(ab_producer_state)
# MMA WARP
if warp_idx == self.mma_warp_id and initial_work_tile_info.is_valid_tile:
tensormap_manager.init_tensormap_from_atom(tma_atom_a, tensormap_a_smem_ptr, self.mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_b, tensormap_b_smem_ptr, self.mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_sfa, tensormap_sfa_smem_ptr, self.mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_sfb, tensormap_sfb_smem_ptr, self.mma_warp_id)
self.tensormap_ab_init_barrier.arrive_and_wait()
self.tmem_alloc_barrier.arrive_and_wait()
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(self.acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
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)
work_tile = 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:
problem_shape_k = work_tile.group_search_result.problem_shape_k
cur_k_tile_cnt = (problem_shape_k + self.cluster_tile_shape_mnk[2] - 1) // self.cluster_tile_shape_mnk[2]
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
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 < cur_k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if is_leader_cta and not is_k_tile_cnt_zero:
acc_pipeline.producer_acquire(acc_producer_state)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(cur_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)
cute.copy(tiled_copy_s2t_sfa, tCsSFA_compact_s2t[s2t_stage_coord], tCtSFA_compact_s2t)
cute.copy(tiled_copy_s2t_sfb, tCsSFB_compact_s2t[s2t_stage_coord], 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[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 < cur_k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if not is_k_tile_cnt_zero:
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
tile_sched.advance_to_next_work()
base_work_tile = tile_sched.get_current_work()
group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
acc_pipeline.producer_tail(acc_producer_state)
# EPILOGUE WARPS
if warp_idx < self.mma_warp_id and initial_work_tile_info.is_valid_tile:
tensormap_manager.init_tensormap_from_atom(tma_atom_c, tensormap_c_smem_ptr, self.epilog_warp_id[0])
if warp_idx == self.epilog_warp_id[0]:
cute.arch.alloc_tmem(self.num_tmem_alloc_cols, tmem_holding_buf, is_two_cta=use_2cta_instrs)
self.tmem_alloc_barrier.arrive_and_wait()
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(self.acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
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)
work_tile = 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)
last_group_idx = cutlass.Int32(-1)
while work_tile.is_valid_tile:
grouped_gemm_cta_tile_info = work_tile.group_search_result
cur_group_idx = grouped_gemm_cta_tile_info.group_idx
cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
is_group_changed = cur_group_idx != last_group_idx
if is_group_changed:
ps = (grouped_gemm_cta_tile_info.problem_shape_m, grouped_gemm_cta_tile_info.problem_shape_n, grouped_gemm_cta_tile_info.problem_shape_k)
real_tensor_c = self.make_tensor_abc_for_tensormap_update(cur_group_idx, self.c_dtype, ps, strides_abc, ptrs_abc, 2)
tensormap_manager.update_tensormap(((real_tensor_c),), ((tma_atom_c),), ((tensormap_c_gmem_ptr),), self.epilog_warp_id[0], (tensormap_c_smem_ptr,))
mma_tile_coord_mnl = (grouped_gemm_cta_tile_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape), grouped_gemm_cta_tile_info.cta_tile_idx_n, 0)
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)]
if not is_k_tile_cnt_zero:
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))
if is_group_changed:
if warp_idx == self.epilog_warp_id[0]:
tensormap_manager.fence_tensormap_update(tensormap_c_gmem_ptr)
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in range(subtile_cnt):
if not is_k_tile_cnt_zero:
cute.copy(tiled_copy_t2r, tTR_tAcc[(None, None, None, subtile_idx)], tTR_rAcc)
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
tRS_rC.store(acc_vec.to(self.c_dtype))
else:
if cutlass.const_expr(self.is_nvfp4_output):
zeros_i8 = cute.make_rmem_tensor(cute.recast_layout(cutlass.Int8.width, self.c_dtype.width, tRS_rC.layout), cutlass.Int8)
zeros_i8.fill(0)
tRS_rC.store(cute.recast_tensor(zeros_i8, self.c_dtype).load())
else:
tRS_rC.fill(0)
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)], tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_c_gmem_ptr, cute.AddressSpace.generic))
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
if not is_k_tile_cnt_zero:
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tile_sched.advance_to_next_work()
base_work_tile = tile_sched.get_current_work()
group_search_result = group_helper.delinearize_z(base_work_tile.tile_idx, problem_sizes_mnkl)
work_tile = GroupedWorkTileInfo(base_work_tile.tile_idx, base_work_tile.is_valid_tile, group_search_result)
last_group_idx = cur_group_idx
if warp_idx == self.epilog_warp_id[0]:
cute.arch.relinquish_tmem_alloc_permit(is_two_cta=use_2cta_instrs)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
if use_2cta_instrs:
cute.arch.mbarrier_arrive(tmem_dealloc_mbar_ptr, cta_rank_in_cluster ^ 1)
cute.arch.mbarrier_wait(tmem_dealloc_mbar_ptr, 0)
cute.arch.dealloc_tmem(acc_tmem_ptr, self.num_tmem_alloc_cols, is_two_cta=use_2cta_instrs)
c_pipeline.producer_tail()
@cute.jit
def make_tensor_abc_for_tensormap_update(self, group_idx: cutlass.Int32, dtype: Type[cutlass.Numeric],
problem_shape_mnk: tuple, strides_abc: cute.Tensor,
tensor_address_abc: cute.Tensor, tensor_index: int):
"""Create tensor for A/B/C tensormap update."""
ptr_i64 = tensor_address_abc[(group_idx, tensor_index)]
tensor_gmem_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
strides_tensor_gmem = strides_abc[(group_idx, tensor_index, None)]
strides_tensor_reg = cute.make_rmem_tensor(cute.make_layout(2), strides_abc.element_type)
cute.autovec_copy(strides_tensor_gmem, strides_tensor_reg)
stride_mn, stride_k = strides_tensor_reg[0], strides_tensor_reg[1]
c1, c0 = cutlass.Int32(1), cutlass.Int32(0)
if cutlass.const_expr(tensor_index == 0):
return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[0], problem_shape_mnk[2], c1), stride=(stride_mn, stride_k, c0)))
elif cutlass.const_expr(tensor_index == 1):
return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[1], problem_shape_mnk[2], c1), stride=(stride_mn, stride_k, c0)))
else:
return cute.make_tensor(tensor_gmem_ptr, cute.make_layout((problem_shape_mnk[0], problem_shape_mnk[1], c1), stride=(stride_mn, stride_k, c0)))
@cute.jit
def make_tensor_sfasfb_for_tensormap_update(self, group_idx: cutlass.Int32, dtype: Type[cutlass.Numeric],
problem_shape_mnk: tuple, tensor_address_sfasfb: cute.Tensor, tensor_index: int):
"""Create tensor for SFA/SFB tensormap update."""
ptr_i64 = tensor_address_sfasfb[(group_idx, tensor_index)]
tensor_gmem_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
c1 = cutlass.Int32(1)
if cutlass.const_expr(tensor_index == 0):
sf_layout = blockscaled_utils.tile_atom_to_shape_SF((problem_shape_mnk[0], problem_shape_mnk[2], c1), self.sf_vec_size)
else:
sf_layout = blockscaled_utils.tile_atom_to_shape_SF((problem_shape_mnk[1], problem_shape_mnk[2], c1), self.sf_vec_size)
return cute.make_tensor(tensor_gmem_ptr, sf_layout)
def mainloop_s2t_copy_and_partition(self, sSF: cute.Tensor, tSF: cute.Tensor) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
"""Create S2T copy for scale factors."""
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]):
"""Create T2R copy for epilogue."""
copy_atom_t2r = sm100_utils.get_tmem_load_op(self.cta_tile_shape_mnk, self.c_layout, self.c_dtype, self.acc_dtype, epi_tile, use_2cta_instrs)
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)])
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
gC_mnl_epi = cute.flat_divide(gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
tTR_rAcc = cute.make_rmem_tensor(tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype)
return tiled_copy_t2r, tTR_tAcc, tTR_rAcc
def epilog_smem_copy_and_partition(self, tiled_copy_t2r: cute.TiledCopy, tTR_rC: cute.Tensor, tidx: cutlass.Int32, sC: cute.Tensor):
"""Create R2S copy for epilogue."""
copy_atom_r2s = sm100_utils.get_smem_store_op(self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r)
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(self, tidx: cutlass.Int32, atom: Union[cute.CopyAtom, cute.TiledCopy],
gC_mnl: cute.Tensor, epi_tile: cute.Tile, sC: cute.Tensor):
"""Create TMA store partition for epilogue."""
gC_epi = cute.flat_divide(gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile)
sC_for_tma = cute.group_modes(sC, 0, 2)
gC_for_tma = cute.group_modes(gC_epi, 0, 2)
bSG_sC, bSG_gC = cpasync.tma_partition(atom, 0, cute.make_layout(1), sC_for_tma, gC_for_tma)
return atom, bSG_sC, bSG_gC
@staticmethod
def _compute_stages(tiled_mma, mma_tiler_mnk, a_dtype, b_dtype, epi_tile, c_dtype, c_layout, sf_dtype, sf_vec_size, smem_capacity, occupancy):
"""Compute pipeline stages."""
num_acc_stage = 1 # Single accumulator stage
num_c_stage = 2
a_smem_layout_1 = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, a_dtype, 1)
b_smem_layout_1 = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, b_dtype, 1)
sfa_smem_layout_1 = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
sfb_smem_layout_1 = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
c_smem_layout_1 = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)
ab_bytes = (cute.size_in_bytes(a_dtype, a_smem_layout_1) + cute.size_in_bytes(b_dtype, b_smem_layout_1) +
cute.size_in_bytes(sf_dtype, sfa_smem_layout_1) + cute.size_in_bytes(sf_dtype, sfb_smem_layout_1))
mbar_bytes = 1024
c_bytes = cute.size_in_bytes(c_dtype, c_smem_layout_1) * num_c_stage
num_ab_stage = (smem_capacity // occupancy - (mbar_bytes + c_bytes)) // ab_bytes
remaining = smem_capacity - occupancy * ab_bytes * num_ab_stage - occupancy * (mbar_bytes + c_bytes)
if remaining > 0:
num_c_stage += remaining // (occupancy * cute.size_in_bytes(c_dtype, c_smem_layout_1))
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
def _compute_grid(total_num_clusters, cluster_shape_mn, max_active_clusters):
"""Compute grid and scheduler params."""
problem_shape = (cluster_shape_mn[0], cluster_shape_mn[1], cutlass.Int32(total_num_clusters))
tile_sched_params = utils.PersistentTileSchedulerParams(problem_shape, (*cluster_shape_mn, 1))
grid = utils.StaticPersistentTileScheduler.get_grid_shape(tile_sched_params, max_active_clusters)
return tile_sched_params, grid
@staticmethod
def _get_mbar_smem_bytes(**kwargs_stages):
return sum(2 * 8 * stage for stage in kwargs_stages.values())
@staticmethod
def is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, c_dtype):
if ab_dtype not in {cutlass.Float4E2M1FN, cutlass.Float8E5M2, cutlass.Float8E4M3FN}: return False
if sf_vec_size not in {16, 32}: return False
if sf_dtype not in {cutlass.Float8E8M0FNU, cutlass.Float8E4M3FN}: return False
if sf_dtype == cutlass.Float8E4M3FN and sf_vec_size == 32: return False
if ab_dtype in {cutlass.Float8E5M2, cutlass.Float8E4M3FN} and sf_vec_size == 16: return False
if c_dtype not in {cutlass.Float32, cutlass.Float16, cutlass.BFloat16, cutlass.Float8E5M2, cutlass.Float8E4M3FN}: return False
return True
@staticmethod
def is_valid_layouts(ab_dtype, c_dtype, a_major, b_major, c_major):
if ab_dtype is cutlass.Float4E2M1FN and not (a_major == "k" and b_major == "k"): return False
return True
@staticmethod
def is_valid_mma_tiler_and_cluster_shape(mma_tiler_mn, cluster_shape_mn):
if mma_tiler_mn[0] not in [128, 256] or mma_tiler_mn[1] not in [128, 256]: return False
if cluster_shape_mn[0] % (2 if mma_tiler_mn[0] == 256 else 1) != 0: return False
is_pow2 = lambda x: x > 0 and (x & (x - 1)) == 0
if (cluster_shape_mn[0] * cluster_shape_mn[1] > 16 or cluster_shape_mn[0] <= 0 or cluster_shape_mn[1] <= 0 or
cluster_shape_mn[0] > 4 or cluster_shape_mn[1] > 4 or not is_pow2(cluster_shape_mn[0]) or not is_pow2(cluster_shape_mn[1])):
return False
return True
@staticmethod
def is_valid_tensor_alignment(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major):
def check_align(dtype, is_mode0_major, shape):
major_idx = 0 if is_mode0_major else 1
return shape[major_idx] % (16 * 8 // dtype.width) == 0
for m, n, k, l in problem_sizes_mnkl:
if not check_align(ab_dtype, a_major == "m", (m, k, l)): return False
if not check_align(ab_dtype, b_major == "n", (n, k, l)): return False
if not check_align(c_dtype, c_major == "m", (m, n, l)): return False
return True
@staticmethod
def can_implement(ab_dtype, sf_dtype, sf_vec_size, c_dtype, mma_tiler_mn, cluster_shape_mn, problem_sizes_mnkl, a_major, b_major, c_major):
K = Sm100GroupedBlockScaledGemmKernel
return (K.is_valid_dtypes_and_scale_factor_vec_size(ab_dtype, sf_dtype, sf_vec_size, c_dtype) and
K.is_valid_layouts(ab_dtype, c_dtype, a_major, b_major, c_major) and
K.is_valid_mma_tiler_and_cluster_shape(mma_tiler_mn, cluster_shape_mn) and
K.is_valid_tensor_alignment(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major))
# HOST FUNCTIONS
def create_tensor_and_stride(l, mode0, mode1, is_mode0_major, dtype, is_dynamic_layout=True):
torch_cpu = cutlass_torch.matrix(l, mode0, mode1, is_mode0_major, cutlass.Float32)
cute_tensor, torch_tensor = cutlass_torch.cute_tensor_like(torch_cpu, dtype, is_dynamic_layout, assumed_align=16)
stride = (1, mode0) if is_mode0_major else (mode1, 1)
return torch_tensor.data_ptr(), torch_tensor, cute_tensor, torch_cpu, stride
def create_tensors_abc_for_all_groups(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major):
ref_tensors, torch_tensors, cute_tensors, strides, ptrs = [], [], [], [], []
for m, n, k, l in problem_sizes_mnkl:
ptr_a, t_a, c_a, r_a, s_a = create_tensor_and_stride(l, m, k, a_major == "m", ab_dtype)
ptr_b, t_b, c_b, r_b, s_b = create_tensor_and_stride(l, n, k, b_major == "n", ab_dtype)
ptr_c, t_c, c_c, r_c, s_c = create_tensor_and_stride(l, m, n, c_major == "m", c_dtype)
ref_tensors.append([r_a, r_b, r_c])
ptrs.append([ptr_a, ptr_b, ptr_c])
torch_tensors.append([t_a, t_b, t_c])
strides.append([s_a, s_b, s_c])
cute_tensors.append((c_a, c_b, c_c))
return ptrs, torch_tensors, cute_tensors, strides, ref_tensors
@cute.jit
def cvt_sf_MKL_to_M32x4xrm_K4xrk_L(sf_ref_tensor: cute.Tensor, sf_mma_tensor: cute.Tensor):
sf_mma_tensor = cute.group_modes(cute.group_modes(sf_mma_tensor, 0, 3), 1, 3)
for i in cutlass.range(cute.size(sf_ref_tensor)):
mkl_coord = sf_ref_tensor.layout.get_hier_coord(i)
sf_mma_tensor[mkl_coord] = sf_ref_tensor[mkl_coord]
def create_scale_factor_tensor(l, mn, k, sf_vec_size, dtype):
def ceil_div(a, b): return (a + b - 1) // b
sf_k = max(1, ceil_div(k, sf_vec_size))
ref_shape, mma_shape = (l, mn, sf_k), (l, ceil_div(mn, 128), ceil_div(sf_k, 4), 32, 4, 4)
ref_f32 = cutlass_torch.create_and_permute_torch_tensor(ref_shape, torch.float32, permute_order=(1,2,0),
init_type=cutlass_torch.TensorInitType.RANDOM, init_config=cutlass_torch.RandomInitConfig(min_val=1, max_val=3))
cute_f32 = cutlass_torch.create_and_permute_torch_tensor(mma_shape, torch.float32, permute_order=(3,4,1,5,2,0),
init_type=cutlass_torch.TensorInitType.RANDOM, init_config=cutlass_torch.RandomInitConfig(min_val=0, max_val=1))
cvt_sf_MKL_to_M32x4xrm_K4xrk_L(from_dlpack(ref_f32), from_dlpack(cute_f32))
cute_f32_gpu = cute_f32.cuda()
ref_f32 = ref_f32.permute(2,0,1).unsqueeze(-1).expand(l,mn,sf_k,sf_vec_size).reshape(l,mn,sf_k*sf_vec_size).permute(1,2,0)[:,:k,:]
cute_tensor, cute_torch = cutlass_torch.cute_tensor_like(cute_f32, dtype, is_dynamic_layout=True, assumed_align=16)
cute_tensor = cutlass_torch.convert_cute_tensor(cute_f32_gpu, cute_tensor, dtype, is_dynamic_layout=True)
return ref_f32, cute_torch.data_ptr(), cute_tensor, cute_torch
def create_tensors_sfasfb_for_all_groups(problem_sizes_mnkl, sf_dtype, sf_vec_size):
ptrs, torch_tensors, cute_tensors, refs = [], [], [], []
for m, n, k, l in problem_sizes_mnkl:
sfa_ref, ptr_sfa, sfa_t, sfa_torch = create_scale_factor_tensor(l, m, k, sf_vec_size, sf_dtype)
sfb_ref, ptr_sfb, sfb_t, sfb_torch = create_scale_factor_tensor(l, n, k, sf_vec_size, sf_dtype)
ptrs.append([ptr_sfa, ptr_sfb])
torch_tensors.append([sfa_torch, sfb_torch])
cute_tensors.append((sfa_t, sfb_t))
refs.append([sfa_ref, sfb_ref])
return ptrs, torch_tensors, cute_tensors, refs
def run(num_groups, problem_sizes_mnkl, host_problem_shape_available, ab_dtype, sf_dtype, sf_vec_size,
c_dtype, a_major, b_major, c_major, mma_tiler_mn, cluster_shape_mn, tolerance=1e-01,
warmup_iterations=0, iterations=1, skip_ref_check=False, **kwargs):
"""Run grouped blockscaled GEMM."""
print(f"Running Blackwell Grouped GEMM: {num_groups} groups, MMA {mma_tiler_mn}, Cluster {cluster_shape_mn}")
if not Sm100GroupedBlockScaledGemmKernel.can_implement(ab_dtype, sf_dtype, sf_vec_size, c_dtype, mma_tiler_mn, cluster_shape_mn, problem_sizes_mnkl, a_major, b_major, c_major):
raise TypeError("Unsupported configuration")
if not torch.cuda.is_available():
raise RuntimeError("GPU required")
torch.manual_seed(2025)
ptrs_abc, torch_abc, cute_abc, strides_abc, ref_abc = create_tensors_abc_for_all_groups(problem_sizes_mnkl, ab_dtype, c_dtype, a_major, b_major, c_major)
ptrs_sfasfb, torch_sfasfb, cute_sfasfb, refs_sfasfb = create_tensors_sfasfb_for_all_groups(problem_sizes_mnkl, sf_dtype, sf_vec_size)
# Initial tensors
alignment = 16
div_ab = 32 if ab_dtype == cutlass.Float4E2M1FN else 16
div_c = 32 if c_dtype == cutlass.Float4E2M1FN else 16
div_sf = 32 if sf_dtype == cutlass.Float4E2M1FN else 16
min_ab = alignment * 8 // ab_dtype.width * ((div_ab + alignment * 8 // ab_dtype.width - 1) // (alignment * 8 // ab_dtype.width))
min_c = alignment * 8 // c_dtype.width * ((div_c + alignment * 8 // c_dtype.width - 1) // (alignment * 8 // c_dtype.width))
min_sf = alignment * 8 // sf_dtype.width * ((div_sf + alignment * 8 // sf_dtype.width - 1) // (alignment * 8 // sf_dtype.width))
init_abc = [create_tensor_and_stride(1, min_ab, min_ab, a_major=="m", ab_dtype)[2],
create_tensor_and_stride(1, min_ab, min_ab, b_major=="n", ab_dtype)[2],
create_tensor_and_stride(1, min_c, min_c, c_major=="m", c_dtype)[2]]
init_sf = [create_tensor_and_stride(1, min_sf, min_sf, a_major=="m", sf_dtype)[2],
create_tensor_and_stride(1, min_sf, min_sf, b_major=="n", sf_dtype)[2]]
hw = cutlass.utils.HardwareInfo()
sm_count = hw.get_max_active_clusters(1)
max_active = hw.get_max_active_clusters(cluster_shape_mn[0] * cluster_shape_mn[1])
tm_shape = (sm_count, Sm100GroupedBlockScaledGemmKernel.num_tensormaps, Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8)
tensor_of_tensormap, _ = cutlass_torch.cute_tensor_like(torch.empty(tm_shape, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False)
gemm = Sm100GroupedBlockScaledGemmKernel(sf_vec_size, mma_tiler_mn, cluster_shape_mn)
td, _ = cutlass_torch.cute_tensor_like(torch.tensor(problem_sizes_mnkl, dtype=torch.int32), cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ts, _ = cutlass_torch.cute_tensor_like(torch.tensor([[[k,1],[k,1],[n,1]] for m,n,k,l in problem_sizes_mnkl], dtype=torch.int32), cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ta, _ = cutlass_torch.cute_tensor_like(torch.tensor(ptrs_abc, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
tf, _ = cutlass_torch.cute_tensor_like(torch.tensor(ptrs_sfasfb, dtype=torch.int64), cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
cta_m = 128 * cluster_shape_mn[0]
cta_n = mma_tiler_mn[1] * cluster_shape_mn[1]
tc = sum(((m+cta_m-1)//cta_m)*((n+cta_n-1)//cta_n) for m,n,_,_ in problem_sizes_mnkl)
if not host_problem_shape_available: tc = max_active
compiled = cute.compile(gemm, init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], num_groups, td, ts, ta, tf, tc, tensor_of_tensormap, max_active, options="--opt-level 2")
if not skip_ref_check:
compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
print("Verifying...")
for i, ((a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (a_t, b_t, c_t), (m, n, k, l)) in enumerate(zip(ref_abc, refs_sfasfb, cute_abc, problem_sizes_mnkl)):
ref = torch.einsum("mkl,nkl->mnl", torch.einsum("mkl,mkl->mkl", a_ref, sfa_ref), torch.einsum("nkl,nkl->nkl", b_ref, sfb_ref))
c_gpu = c_ref.cuda()
cute.testing.convert(c_t, from_dlpack(c_gpu, assumed_align=16).mark_layout_dynamic(leading_dim=(1 if c_major=="n" else 0)))
c_cpu = c_gpu.cpu()
if c_dtype in (cutlass.Float32, cutlass.Float16, cutlass.BFloat16):
torch.testing.assert_close(c_cpu, ref, atol=tolerance, rtol=1e-02)
elif c_dtype in (cutlass.Float8E5M2, cutlass.Float8E4M3FN):
ref_f8_ = torch.empty(*(l,m,n), dtype=torch.uint8, device="cuda").permute(1,2,0)
ref_f8 = from_dlpack(ref_f8_, assumed_align=16).mark_layout_dynamic(leading_dim=1)
ref_f8.element_type = c_dtype
ref_dev = ref.permute(2,0,1).contiguous().permute(1,2,0).cuda()
ref_t = from_dlpack(ref_dev, assumed_align=16).mark_layout_dynamic(leading_dim=1)
cute.testing.convert(ref_t, ref_f8)
cute.testing.convert(ref_f8, ref_t)
torch.testing.assert_close(c_cpu, ref_dev.cpu(), atol=tolerance, rtol=1e-02)
for _ in range(warmup_iterations):
compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
torch.cuda.synchronize()
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(iterations):
compiled(init_abc[0], init_abc[1], init_abc[2], init_sf[0], init_sf[1], td, ts, ta, tf, tensor_of_tensormap)
end.record()
torch.cuda.synchronize()
exec_time = start.elapsed_time(end) * 1000 / iterations
fmas = sum(m*n*k for m,n,k,_ in problem_sizes_mnkl)
gflops = 2 * fmas / 1e9 / (exec_time / 1e6)
print(f"Runtime: {exec_time/1000:.3f}ms, GFLOPS: {gflops:.1f}")
return exec_time
if __name__ == "__main__":
def parse_ints(s): return tuple(int(x.strip()) for x in s.split(","))
def parse_tuples(s):
tuples = s.strip("()").split("),(")
return [tuple(int(x.strip()) for x in t.split(",")) for t in tuples]
parser = argparse.ArgumentParser(description="Grouped GEMM on Blackwell")
parser.add_argument("--num_groups", type=int, default=2)
parser.add_argument("--problem_sizes_mnkl", type=parse_tuples, default=((128,128,128,1),(128,128,128,1)))
parser.add_argument("--mma_tiler_mn", type=parse_ints, default=(128,128))
parser.add_argument("--host_problem_shape_available", action="store_true")
parser.add_argument("--cluster_shape_mn", type=parse_ints, default=(1,1))
parser.add_argument("--ab_dtype", type=cutlass.dtype, default=cutlass.Float4E2M1FN)
parser.add_argument("--sf_dtype", type=cutlass.dtype, default=cutlass.Float8E8M0FNU)
parser.add_argument("--sf_vec_size", type=int, default=16)
parser.add_argument("--c_dtype", type=cutlass.dtype, default=cutlass.Float16)
parser.add_argument("--a_major", choices=["k","m"], default="k")
parser.add_argument("--b_major", choices=["k","n"], default="k")
parser.add_argument("--c_major", choices=["n","m"], default="n")
parser.add_argument("--tolerance", type=float, default=1e-01)
parser.add_argument("--warmup_iterations", type=int, default=0)
parser.add_argument("--iterations", type=int, default=1)
parser.add_argument("--skip_ref_check", action="store_true")
args = parser.parse_args()
if len(args.problem_sizes_mnkl) != args.num_groups:
parser.error("problem_sizes_mnkl must contain num_groups tuples")
for _, _, _, l in args.problem_sizes_mnkl:
if l != 1: parser.error("l must be 1")
run(args.num_groups, args.problem_sizes_mnkl, args.host_problem_shape_available, args.ab_dtype,
args.sf_dtype, args.sf_vec_size, args.c_dtype, args.a_major, args.b_major, args.c_major,
args.mma_tiler_mn, args.cluster_shape_mn, args.tolerance, args.warmup_iterations,
args.iterations, args.skip_ref_check)
print("PASS")
# AUTO-TESTING ENTRY POINT
try:
from task import input_t, output_t
except ImportError:
input_t = Tuple[List[Tuple[torch.Tensor, torch.Tensor, torch.Tensor]], List[Tuple[torch.Tensor, torch.Tensor]], List[Tuple[torch.Tensor, torch.Tensor]], List[Tuple[int, int, int, int]]]
output_t = List[torch.Tensor]
_base_configs = {}
_cache, _ptr_cache, _last_call = {}, {}, None
_single_kernel_cache = {} # Cache for single-group kernels
_graph_cache = {} # Cache for CUDA graphs
def _sort_groups_by_k_desc(ps, at, sr):
"""Sort groups by K dimension in descending order."""
indexed = [(i, ps[i][2]) for i in range(len(ps))]
sorted_indexed = sorted(indexed, key=lambda x: -x[1]) # Descending K
sorted_indices = [x[0] for x in sorted_indexed]
inverse_indices = [0] * len(sorted_indices)
for new_idx, old_idx in enumerate(sorted_indices):
inverse_indices[old_idx] = new_idx
return sorted_indices, inverse_indices
def _sort_groups_by_work(ps, at, sr):
"""Sort groups by total work (M*N*K) in descending order for better load balancing."""
indexed = [(i, ps[i][0] * ps[i][1] * ps[i][2]) for i in range(len(ps))]
sorted_indexed = sorted(indexed, key=lambda x: -x[1]) # Descending order
sorted_indices = [x[0] for x in sorted_indexed]
inverse_indices = [0] * len(sorted_indices)
for new_idx, old_idx in enumerate(sorted_indices):
inverse_indices[old_idx] = new_idx
return sorted_indices, inverse_indices
def _sort_groups_by_nk(ps, at, sr):
indexed = [(i, (ps[i][1], ps[i][2])) for i in range(len(ps))]
sorted_indexed = sorted(indexed, key=lambda x: x[1])
sorted_indices = [x[0] for x in sorted_indexed]
inverse_indices = [0] * len(sorted_indices)
for new_idx, old_idx in enumerate(sorted_indices):
inverse_indices[old_idx] = new_idx
return sorted_indices, inverse_indices
def _get_base(config_key):
"""Get or create base tensors for a given configuration."""
global _base_configs
MMA_M, MMA_N, CLUSTER_M, CLUSTER_N = config_key
if config_key not in _base_configs:
ab, sf, c = cutlass.Float4E2M1FN, cutlass.Float8E4M3FN, cutlass.Float16
abc = [create_tensor_and_stride(1,64,64,False,ab)[2], create_tensor_and_stride(1,64,64,False,ab)[2], create_tensor_and_stride(1,16,16,False,c)[2]]
sft = [create_tensor_and_stride(1,16,16,False,sf)[2], create_tensor_and_stride(1,16,16,False,sf)[2]]
hw = cutlass.utils.HardwareInfo()
CLUSTER_SIZE = CLUSTER_M * CLUSTER_N
sm = hw.get_max_active_clusters(CLUSTER_SIZE)
tm = torch.zeros((sm*CLUSTER_SIZE,5,16), dtype=torch.int64, device="cuda")
tmap, _ = cutlass_torch.cute_tensor_like(tm, cutlass.Int64, is_dynamic_layout=False)
_base_configs[config_key] = {"abc": abc, "sf": sft, "tmap": tmap, "max": sm, "cluster_size": CLUSTER_SIZE}
return _base_configs[config_key]
def _get_single_group_kernel(m, n, k, config_key):
"""Get or compile a single-group kernel for the given problem size."""
global _single_kernel_cache
MMA_M, MMA_N, CLUSTER_M, CLUSTER_N = config_key
CLUSTER_TILE_M = 128 * CLUSTER_M
CLUSTER_TILE_N = MMA_N * CLUSTER_N
# Cache key based on tile counts (not exact size) for reuse
tm = (m + CLUSTER_TILE_M - 1) // CLUSTER_TILE_M
tn = (n + CLUSTER_TILE_N - 1) // CLUSTER_TILE_N
tk = k
cache_key = (tm, tn, tk, config_key)
if cache_key not in _single_kernel_cache:
b = _get_base(config_key)
ps = [(m, n, k, 1)]
st = [[[k,1],[k,1],[n,1]]]
tc = tm * tn
d = torch.tensor(ps, dtype=torch.int32, device="cuda")
s = torch.tensor(st, dtype=torch.int32, device="cuda")
pa = torch.zeros((1, 3), dtype=torch.int64, device="cuda")
pf = torch.zeros((1, 2), dtype=torch.int64, device="cuda")
td, _ = cutlass_torch.cute_tensor_like(d, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ts, _ = cutlass_torch.cute_tensor_like(s, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ta, ta_torch = cutlass_torch.cute_tensor_like(pa, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
tf, tf_torch = cutlass_torch.cute_tensor_like(pf, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
g = Sm100GroupedBlockScaledGemmKernel(16, (MMA_M, MMA_N), (CLUSTER_M, CLUSTER_N), occupancy=1)
kernel = cute.compile(g, b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1],
1, td, ts, ta, tf, tc, b["tmap"], b["max"], options="--opt-level 3")
_single_kernel_cache[cache_key] = {
"kernel": kernel, "td": td, "ts": ts, "ta": ta, "tf": tf,
"ta_torch": ta_torch, "tf_torch": tf_torch, "base": b
}
return _single_kernel_cache[cache_key]
def _run_separate_kernels(at, sr, ps, config_key):
"""Run separate single-group kernels for each group."""
ng = len(ps)
outs = []
for i in range(ng):
m, n, k, l = ps[i]
a, b_, c = at[i]
sa, sb = sr[i]
r = _get_single_group_kernel(m, n, k, config_key)
# Update pointers
r["ta_torch"][0, 0] = a.data_ptr()
r["ta_torch"][0, 1] = b_.data_ptr()
r["ta_torch"][0, 2] = c.data_ptr()
r["tf_torch"][0, 0] = sa.data_ptr()
r["tf_torch"][0, 1] = sb.data_ptr()
# Run kernel for this group
b = r["base"]
r["kernel"](b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1],
r["td"], r["ts"], r["ta"], r["tf"], b["tmap"])
outs.append(c)
return outs
def custom_kernel(data: input_t) -> output_t:
global _cache, _ptr_cache, _last_call
at, _, sr, ps = data
ng = len(ps)
# Default config
MMA_M, MMA_N, CLUSTER_M, CLUSTER_N, OCCUPANCY = 128, 128, 1, 1, 1
CLUSTER_TILE_M, CLUSTER_TILE_N = 128 * CLUSTER_M, MMA_N * CLUSTER_N
config_key = (MMA_M, MMA_N, CLUSTER_M, CLUSTER_N)
if _last_call is not None:
last_ng, last_at, last_sr, cached = _last_call
if last_ng == ng and last_at is at and last_sr is sr:
kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap = cached["args"]
kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
return cached["outs"]
b = _get_base(config_key)
ptr_at = tuple((a.data_ptr(), b_.data_ptr(), c.data_ptr()) for (a, b_, c) in at)
ptr_sr = tuple((sa.data_ptr(), sb.data_ptr()) for (sa, sb) in sr)
ptr_key = (ng, ptr_at, ptr_sr, tuple(tuple(x) for x in ps), config_key)
if ptr_key in _ptr_cache:
r, outs, cached_args = _ptr_cache[ptr_key]
kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap = cached_args
kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
_last_call = (ng, at, sr, {"args": cached_args, "outs": outs})
return outs
# Keep groups in original order - let the scheduler handle distribution
sorted_indices = list(range(ng))
ps_sorted = [ps[i] for i in sorted_indices]
at_sorted = [at[i] for i in sorted_indices]
sr_sorted = [sr[i] for i in sorted_indices]
tc = sum(((m+CLUSTER_TILE_M-1)//CLUSTER_TILE_M)*((n+CLUSTER_TILE_N-1)//CLUSTER_TILE_N) for m,n,_,_ in ps_sorted)
key = (ng, tc, tuple(tuple(x) for x in ps_sorted), config_key)
ap = [[a.data_ptr(), b_.data_ptr(), c.data_ptr()] for (a, b_, c) in at_sorted]
sp = [[sa.data_ptr(), sb.data_ptr()] for (sa, sb) in sr_sorted]
st = [[[k,1],[k,1],[n,1]] for m,n,k,l in ps_sorted]
if key not in _cache:
d = torch.tensor(ps_sorted, dtype=torch.int32, device="cuda")
s = torch.tensor(st, dtype=torch.int32, device="cuda")
pa = torch.tensor(ap, dtype=torch.int64, device="cuda")
pf = torch.tensor(sp, dtype=torch.int64, device="cuda")
td, _ = cutlass_torch.cute_tensor_like(d, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ts, _ = cutlass_torch.cute_tensor_like(s, cutlass.Int32, is_dynamic_layout=False, assumed_align=16)
ta, ta_torch = cutlass_torch.cute_tensor_like(pa, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
tf, tf_torch = cutlass_torch.cute_tensor_like(pf, cutlass.Int64, is_dynamic_layout=False, assumed_align=16)
g = Sm100GroupedBlockScaledGemmKernel(16, (MMA_M, MMA_N), (CLUSTER_M, CLUSTER_N), occupancy=1, max_ab_stages=0)
k = cute.compile(g, b["abc"][0], b["abc"][1], b["abc"][2], b["sf"][0], b["sf"][1], ng, td, ts, ta, tf, tc, b["tmap"], b["max"], options="--opt-level 3")
_cache[key] = {"k": k, "td": td, "ts": ts, "ta": ta, "tf": tf, "ta_torch": ta_torch, "tf_torch": tf_torch}
r = _cache[key]
if "ap_cpu" not in r:
r["ap_cpu"] = torch.zeros(ng*3, dtype=torch.int64, pin_memory=True)
r["sp_cpu"] = torch.zeros(ng*2, dtype=torch.int64, pin_memory=True)
# Direct copy from pinned CPU to cute tensor's underlying torch tensor
ap_cpu, sp_cpu = r["ap_cpu"], r["sp_cpu"]
idx = 0
for (a_, b_, c_) in at_sorted:
ap_cpu[idx], ap_cpu[idx+1], ap_cpu[idx+2] = a_.data_ptr(), b_.data_ptr(), c_.data_ptr()
idx += 3
idx = 0
for (sa, sb) in sr_sorted:
sp_cpu[idx], sp_cpu[idx+1] = sa.data_ptr(), sb.data_ptr()
idx += 2
# Single copy: pinned CPU -> GPU cute tensor
r["ta_torch"].view(-1).copy_(ap_cpu, non_blocking=True)
r["tf_torch"].view(-1).copy_(sp_cpu, non_blocking=True)
outs = [at[i][2] for i in range(ng)]
kernel = r["k"]
abc0, abc1, abc2 = b["abc"][0], b["abc"][1], b["abc"][2]
sf0, sf1 = b["sf"][0], b["sf"][1]
td, ts, ta, tf, tmap = r["td"], r["ts"], r["ta"], r["tf"], b["tmap"]
# Store args for caching
cached_args = (kernel, abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
_ptr_cache[ptr_key] = (r, outs, cached_args)
_last_call = (ng, at, sr, {"args": cached_args, "outs": outs})
# Direct kernel call
kernel(abc0, abc1, abc2, sf0, sf1, td, ts, ta, tf, tmap)
return outs
print("main")
print("main end")
scrolls · 1091 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 403686.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON