submission 155565
rcmalli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1725 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-155565?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:dc2ccaa6a870fcd09ed337220cc85a734a5d44e49cf2f2ebae2251c40a42d29b
license declaredunknown
license concludedunknown
authorsrcmalli
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Custom NVFP4 block-scaled GEMM kernel implementation using CuTe-DSL.fused-epilogue
- Warp specialization: TMA load (warp 5), MMA (warp 4), Epilogue (warps 0-3)mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(persistent-kernel
This implements a persistent warp-specialized kernel for B200 (SM100) targetingshared-memory
self.smem_capacity = 228 * 1024 # 233472 bytestcgen05
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.py1725 lines
"""
Custom NVFP4 block-scaled GEMM kernel implementation using CuTe-DSL.
This implements a persistent warp-specialized kernel for B200 (SM100) targeting
the NVFP4 block-scaled GEMM competition.
Key features:
- Warp specialization: TMA load (warp 5), MMA (warp 4), Epilogue (warps 0-3)
- Block-scaled FP4 with FP8 scale factors (sf_vec_size=16)
- TMA multicast for B matrix across cluster
- Software pipelining for A/B/SF data
"""
from dataclasses import dataclass
from typing import Optional, Tuple, Type, Union
import os
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
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 make_ptr
from task import input_t, output_t
@dataclass(frozen=True)
class SchedulerParams:
tiles_m: int
tiles_n: int
k_tile_cnt: int
@property
def tiles_mn(self) -> int:
return self.tiles_m * self.tiles_n
def total_work_items(self, l: int) -> int:
return self.tiles_mn * l
@classmethod
def from_shape(
cls,
m: int,
n: int,
k_tile_cnt: int,
tile_m: int,
tile_n: int,
):
tiles_m = (m + tile_m - 1) // tile_m
tiles_n = (n + tile_n - 1) // tile_n
return cls(tiles_m, tiles_n, k_tile_cnt)
class Sm100BlockScaledGemmKernel:
"""
Block-scaled FP4 GEMM kernel for SM100 (Blackwell).
Configuration:
- mma_tiler_mn: (128, 128) for M=128 cases
- cluster_shape_mn: (1, 2) for B-matrix multicast
- sf_vec_size: 16 (standard for NVFP4)
"""
def __init__(
self,
sf_vec_size: int = 16,
mma_tiler_mn: Tuple[int, int] = (128, 128),
cluster_shape_mn: Tuple[int, int] = (1, 2),
num_ab_stage: Optional[int] = None,
num_c_stage: Optional[int] = None,
ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN,
sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN,
c_dtype: Type[cutlass.Numeric] = cutlass.Float16,
use_persistent_scheduler: bool = True,
):
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) # K dimension set later
self.use_persistent_scheduler = use_persistent_scheduler
self.num_ab_stage_override = num_ab_stage
self.num_c_stage_override = num_c_stage
# Store data types as instance variables (needed for JIT compilation)
self.a_dtype = ab_dtype
self.b_dtype = ab_dtype
self.sf_dtype = sf_dtype
self.c_dtype = c_dtype
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.occupancy = 1
# Warp specialization IDs
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)
)
# Named barriers
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=self.threads_per_cta,
)
# Keep all warps/roles in lockstep across work items (ported from dev_v6 Strm-K loop)
self.work_loop_barrier = pipeline.NamedBarrier(
barrier_id=3,
num_threads=self.threads_per_cta,
)
# SM100 has 228KB shared memory - hardware constant
self.smem_capacity = 228 * 1024 # 233472 bytes
SM100_TMEM_CAPACITY_COLUMNS = 512
self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS
def _setup_attributes(self):
"""Configure kernel attributes based on input tensor properties."""
# MMA instruction shape
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),
)
# Create tiled MMA for block-scaled operation
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,
)
# Compute MMA/cluster/tile shapes
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.cta_tile_shape_mnk_sfb = (
self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler_sfb[1],
self.mma_tiler_sfb[2],
)
# Cluster layout
self.cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma.thr_id.shape,),
)
self.cluster_layout_sfb_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma_sfb.thr_id.shape,),
)
# Multicast configuration
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
# Epilogue tile
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.epi_tile_n = cute.size(self.epi_tile[1])
# Stage counts
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,
)
if self.num_ab_stage_override is not None and self.num_ab_stage_override > 0:
self.num_ab_stage = max(2, min(8, int(self.num_ab_stage_override)))
if self.num_c_stage_override is not None and self.num_c_stage_override > 0:
self.num_c_stage = max(1, min(4, int(self.num_c_stage_override)))
# SMEM layouts
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,
)
# Accumulator configuration
self.overlapping_accum = self.num_acc_stage == 1
# TMEM column counts
sf_atom_mn = 32
self.num_sfa_tmem_cols = (
self.cta_tile_shape_mnk[0] // sf_atom_mn
) * mma_inst_tile_k
self.num_sfb_tmem_cols = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
) * mma_inst_tile_k
self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
self.num_accumulator_tmem_cols = (
self.cta_tile_shape_mnk[1] * self.num_acc_stage
if not self.overlapping_accum
else self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols
)
self.iter_acc_early_release_in_epilogue = (
self.num_sf_tmem_cols // self.epi_tile_n
)
def _compute_stages(
self,
tiled_mma,
mma_tiler,
a_dtype,
b_dtype,
epi_tile,
c_dtype,
c_layout,
sf_dtype,
sf_vec_size,
smem_capacity,
occupancy,
):
"""Compute pipeline stage counts based on SMEM capacity."""
# Compute SMEM footprints
a_smem_layout = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler, a_dtype, 1)
b_smem_layout = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler, b_dtype, 1)
sfa_smem_layout = blockscaled_utils.make_smem_layout_sfa(
tiled_mma, mma_tiler, sf_vec_size, 1
)
sfb_smem_layout = blockscaled_utils.make_smem_layout_sfb(
tiled_mma, mma_tiler, sf_vec_size, 1
)
c_smem_layout = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)
a_bytes = cute.size_in_bytes(a_dtype, a_smem_layout.outer)
b_bytes = cute.size_in_bytes(b_dtype, b_smem_layout.outer)
sfa_bytes = cute.size_in_bytes(sf_dtype, sfa_smem_layout)
sfb_bytes = cute.size_in_bytes(sf_dtype, sfb_smem_layout)
c_bytes = cute.size_in_bytes(c_dtype, c_smem_layout.outer)
ab_sf_bytes_per_stage = a_bytes + b_bytes + sfa_bytes + sfb_bytes
# Force num_acc_stage=1 to enable overlapping_accum (accumulator ping-pong)
# This allows epilogue to overlap with next tile's MMA via early release
num_acc_stage = 1
# Compute C stages (2 for double buffering)
num_c_stage = 2
# Fixed overhead for barriers, etc.
barrier_bytes = 1024
# Available SMEM for A/B/SF stages
available = smem_capacity // occupancy - c_bytes * num_c_stage - barrier_bytes
num_ab_stage = max(2, min(8, available // ab_sf_bytes_per_stage))
return num_acc_stage, num_ab_stage, num_c_stage
def _compute_grid(
self, c_tensor, cta_tile_shape, cluster_shape_mn, max_active_clusters
):
"""Compute grid dimensions for kernel launch.
Uses CuTe's zipped_divide to compute the number of CTAs needed,
then uses StaticPersistentTileScheduler.get_grid_shape for the grid.
"""
# Use zipped_divide to get the number of CTAs per dimension
c_shape = cute.slice_(cta_tile_shape, (None, None, 0))
gc = cute.zipped_divide(c_tensor, tiler=c_shape)
num_ctas_mnl = gc[(0, (None, None, None))].shape
# PersistentTileSchedulerParams expects cluster_shape_mnk with K=1
cluster_shape_mnk = (*cluster_shape_mn, 1)
tile_sched_params = utils.PersistentTileSchedulerParams(
num_ctas_mnl, cluster_shape_mnk
)
# Use StaticPersistentTileScheduler to compute the grid shape
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
return tile_sched_params, grid
@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,
shape_mnkl: Tuple[int, int, int, int],
max_active_clusters: cutlass.Constexpr = 132,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
"""Execute the block-scaled GEMM operation.
Args:
a_ptr: Pointer to A matrix data (M x K x L, K-major, FP4)
b_ptr: Pointer to B matrix data (N x K x L, K-major, FP4)
sfa_ptr: Pointer to scale factors for A (permuted layout, FP8)
sfb_ptr: Pointer to scale factors for B (permuted layout, FP8)
c_ptr: Pointer to C matrix data (M x N x L, N-major, FP16)
shape_mnkl: Tuple of (M, N, K, L) dimensions
max_active_clusters: Maximum active clusters for persistent scheduling
epilogue_op: Optional epilogue operation
"""
m, n, k, l = shape_mnkl
# Create tensors from pointers with proper layouts
# A is (M, K, L) K-major
a_layout = cute.make_layout((m, k, l), stride=(k, 1, m * k))
a_tensor = cute.make_tensor(a_ptr, a_layout)
# B is (N, K, L) K-major
b_layout = cute.make_layout((n, k, l), stride=(k, 1, n * k))
b_tensor = cute.make_tensor(b_ptr, b_layout)
# C is (M, N, L) row-major (N-stride=1, M-stride=N)
c_layout = cute.make_layout((m, n, l), stride=(n, 1, m * n))
c_tensor = cute.make_tensor(c_ptr, c_layout)
# Get major modes from layouts
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)
# Type check
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")
# Setup kernel attributes
self._setup_attributes()
# Setup scale factor tensor layouts from pointers
# Scale factors use the tile_atom_to_shape_SF layout based on A/B shapes
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)
# Create tiled MMAs
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)
# Setup TMA atoms 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,
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,
)
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
# TMA store for C
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,
)
# ---------------------------------------------------------------------
# Deterministic persistent scheduling
#
# Launch a single logical cluster in (x,y) and use grid.z as a persistent
# work queue length. All warps participate in a shared work loop and
# synchronize once per work item.
# ---------------------------------------------------------------------
cluster_m, cluster_n = self.cluster_shape_mn
v_ctas_per_tile = cute.size(tiled_mma.thr_id.shape)
scheduler_params = SchedulerParams.from_shape(
m=c_tensor.shape[0],
n=c_tensor.shape[1],
k_tile_cnt=k // self.mma_tiler[2],
tile_m=self.mma_tiler[0] * cluster_m,
tile_n=self.mma_tiler[1] * cluster_n,
)
tiles_m = scheduler_params.tiles_m
tiles_n = scheduler_params.tiles_n
total_work_items = scheduler_params.total_work_items(l)
# When persistent scheduling is disabled, run one work item per cluster.
if not self.use_persistent_scheduler:
max_active_clusters = total_work_items
grid_z = min(total_work_items, max_active_clusters)
grid = (cluster_m * v_ctas_per_tile, cluster_n, grid_z)
self.buffer_align_bytes = 1024
# Define shared storage
@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
self.kernel(
tiled_mma,
tiled_mma_sfb,
tma_atom_a,
tma_tensor_a,
tma_atom_b,
tma_tensor_b,
tma_atom_sfa,
tma_tensor_sfa,
tma_atom_sfb,
tma_tensor_sfb,
tma_atom_c,
tma_tensor_c,
c_tensor,
self.cluster_layout_vmnk,
self.cluster_layout_sfb_vmnk,
self.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,
tiles_m,
tiles_n,
l,
total_work_items,
epilogue_op,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(
self.cluster_shape_mn[0] * v_ctas_per_tile,
self.cluster_shape_mn[1],
1,
),
min_blocks_per_mp=1,
)
@cute.kernel
def kernel(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a: cute.CopyAtom,
mA_mkl: cute.Tensor,
tma_atom_b: cute.CopyAtom,
mB_nkl: cute.Tensor,
tma_atom_sfa: cute.CopyAtom,
mSFA_mkl: cute.Tensor,
tma_atom_sfb: cute.CopyAtom,
mSFB_nkl: cute.Tensor,
tma_atom_c: cute.CopyAtom,
mC_mnl: cute.Tensor,
mC_raw_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,
tiles_m: cutlass.Int32,
tiles_n: cutlass.Int32,
tiles_l: cutlass.Int32,
total_work_items: cutlass.Int32,
epilogue_op: cutlass.Constexpr,
):
"""
GPU device kernel for block-scaled GEMM.
Kernel argument shapes (logical, as seen by the harness / custom_kernel):
- A: `a` is (M, K_packed, L) where `K_packed = K // 2` due to FP4 packing.
This kernel treats it as logical (M, K, L) by doubling K when needed.
- B: `b` is (N, K_packed, L) with the same packing along K.
- SFA (permuted): `sfa_permuted` is a CuTe/TMA-friendly view of scale factors for A.
Reference `sfa` is (M, K//16, L); permuted view is typically shaped like
(32, 4, rest_m, 4, rest_k, L) (exact factorization depends on tiling).
- SFB (permuted): `sfb_permuted` is analogous for B; reference `sfb` is (N, K//16, L),
permuted view is typically (32, 4, rest_n, 4, rest_k, L).
- C: `c` is (M, N, L) fp16; the harness uses L as the fastest-varying dimension.
"""
# CuTe arch intrinsics:
# - `cute.arch.warp_idx()` gives the warp id within the CTA (0..).
# - `cute.arch.make_warp_uniform(x)` forces a warp-uniform value (avoid divergence pitfalls).
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# TMA descriptor prefetch:
# - `cpasync.prefetch_descriptor(tma_atom)` hides descriptor fetch latency before first TMA issue.
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
# CTA/thread coordinates:
# - `cute.arch.block_idx()` -> (blockIdx.x, blockIdx.y, blockIdx.z).
# - `cute.arch.block_idx_in_cluster()` -> CTA rank within the hardware cluster.
bidx, bidy, bidz = cute.arch.block_idx()
v_ctas_per_tile = cute.size(tiled_mma.thr_id.shape)
mma_tile_coord_v = bidx % v_ctas_per_tile
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
# Cluster layout decode:
# - `layout.get_flat_coord(rank)` maps CTA-rank -> multi-dim coords used by multicast + partitions.
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()
# For our launch, grid.x/y enumerate CTAs within a cluster.
cluster_n_coord = cute.arch.make_warp_uniform(bidy)
# Shared memory allocation helper:
# - `utils.SmemAllocator()` + `allocate(self.shared_storage)` builds a typed SMEM backing store.
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
# Pipelines:
# - `PipelineTmaUmma` coordinates TMA producer(s) and UMMA consumer(s) via mbarrier arrays.
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,
defer_sync=True,
)
# Accumulator pipeline:
# - `PipelineUmmaAsync` synchronizes "accumulator stage" ownership across MMA and epilogue roles.
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,
defer_sync=True,
)
# TMEM allocator:
# - Provides per-CTA allocations of Blackwell tensor memory columns for accumulators + SF tiles.
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,
)
# Cluster arrive after barrier init
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)
# TMA multicast masks:
# - `cpasync.create_tma_multicast_mask(...)` returns the mask consumed by `cute.copy(..., mcast_mask=...)`.
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
)
# Global -> per-CTA tiles:
# - `cute.local_tile(tensor, slice, coords)` forms a logical tile view for this CTA over GMEM 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_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)
)
# Tile sizes:
# - `cute.size(tensor, mode=[...])` returns the extent of specific tensor mode(s) (here: K-tiles).
k_tile_cnt = cute.size(gA_mkl, mode=[3])
# 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:
# - `cpasync.tma_partition(tma_atom, cta_coord, cta_layout, smem_tile, gmem_tile)`
# returns (SMEM view, GMEM view) for issuing TMA copies stage-by-stage.
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),
)
# TMA partitions for B
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),
)
# TMA partitions for SFA
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cpasync.tma_partition(
tma_atom_sfa,
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA, 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
# TMA partitions for SFB
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 = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
# MMA fragment tensors
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])
if cutlass.const_expr(self.overlapping_accum):
num_acc_stage_overlapped = 2
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, num_acc_stage_overlapped)
)
tCtAcc_fake = cute.make_tensor(
tCtAcc_fake.iterator,
cute.make_layout(
tCtAcc_fake.shape,
stride=(
tCtAcc_fake.stride[0],
tCtAcc_fake.stride[1],
tCtAcc_fake.stride[2],
(256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1],
),
),
)
else:
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
# Wait for cluster init
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
# ---------------------------------------------------------------------
# Deterministic persistent work loop (dev_v6 pattern)
# ---------------------------------------------------------------------
grid_dim_z = cute.arch.grid_dim()[2]
grid_dim_z = cute.arch.make_warp_uniform(grid_dim_z)
base_work_idx = cute.arch.make_warp_uniform(bidz)
tiles_mn = tiles_m * tiles_n
num_work = cute.ceil_div(total_work_items - base_work_idx, grid_dim_z)
# Per-CTA pipeline states (must be defined outside dynamic control flow)
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
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
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
# Allocate + retrieve TMEM pointers (once) (ported from dev_v6)
tmem.allocate(self.num_tmem_alloc_cols)
self.tmem_alloc_barrier.arrive_and_wait()
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 + self.num_accumulator_tmem_cols,
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 + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
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)
# Epilogue copy partitions (once)
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_gC_base,
(tTR_rAcc_0, tTR_rAcc_1),
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc_0.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(tma_atom_c, tCgC, epi_tile, sC)
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,
)
# ---------------------------------------------------------------------
# Main work loop
# ---------------------------------------------------------------------
for work_iter in cutlass.range(num_work, unroll=1):
work_idx = base_work_idx + work_iter * grid_dim_z
# No-split schedule: work_idx -> (m,n,l)
work_m_idx = work_idx % tiles_m
work_n_idx = (work_idx // tiles_m) % tiles_n
work_l_idx = work_idx // tiles_mn
cluster_tile_coord_mnl = (work_m_idx, work_n_idx, work_l_idx)
mma_tile_coord_mnl = (
work_m_idx,
work_n_idx * self.cluster_shape_mn[1] + cluster_n_coord,
work_l_idx,
)
# -------------------------------------------------------------
# TMA LOAD WARP (warp 5)
# -------------------------------------------------------------
if warp_idx == self.tma_warp_id:
tAgA_slice = tAgA[
(None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
]
tBgB_slice = tBgB[
(None, mma_tile_coord_mnl[1], None, cluster_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA[
(None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])
]
slice_n = cluster_tile_coord_mnl[1]
if cutlass.const_expr(self.cluster_shape_mn[1] > 1):
# For clustered N, SFB needs per-CTA N tile selection (tile_n chunks).
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = slice_n // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, cluster_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_iter in cutlass.range(0, k_tile_cnt, 1, unroll=1): # noqa: B007
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
k_tile = ab_producer_state.count
# TMA loads:
# - `cute.copy(tma_atom_*, gmem_tile, smem_tile, tma_bar_ptr=..., mcast_mask=...)`
# issues a TMA transaction and signals the stage mbarrier when complete.
cute.copy(
tma_atom_a,
tAgA_slice[(None, k_tile)],
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, k_tile)],
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, k_tile)],
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, k_tile)],
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
)
# -------------------------------------------------------------
# MMA WARP (warp 4)
# -------------------------------------------------------------
if warp_idx == self.mma_warp_id:
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_producer_state.phase ^ 1
else:
acc_stage_index = acc_producer_state.index
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_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
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ 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
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_iter in range(k_tile_cnt): # noqa: B007
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
)
# MMA:
# - `cute.gemm(tiled_mma, D, A, B, C)` emits tcgen05 MMA into TMEM accumulator `D`.
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()
# -------------------------------------------------------------
# EPILOGUE WARPS (warps 0-3)
# -------------------------------------------------------------
if warp_idx < self.mma_warp_id:
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
else:
acc_stage_index = acc_consumer_state.index
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
# Mode grouping:
# - `cute.group_modes(t, start, end)` merges modes to produce a simpler iteration space.
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
tTR_gC_tile = tTR_gC_base[
(
None,
None,
None,
None,
None,
mma_tile_coord_mnl[0],
mma_tile_coord_mnl[1],
mma_tile_coord_mnl[2],
)
]
tTR_gC = cute.group_modes(tTR_gC_tile, 3, cute.rank(tTR_gC_tile))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = 0
if warp_idx == self.epilog_warp_id[0] and subtile_cnt > 0:
c_pipeline.producer_acquire()
if subtile_cnt > 0:
tTR_tAcc_mn_first = tTR_tAcc[(None, None, None, 0)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn_first, tTR_rAcc_0)
for subtile_idx in range(subtile_cnt):
curr_buf = tTR_rAcc_0 if subtile_idx % 2 == 0 else tTR_rAcc_1
next_buf = tTR_rAcc_1 if subtile_idx % 2 == 0 else tTR_rAcc_0
if subtile_idx + 1 < subtile_cnt:
tTR_tAcc_mn_next = tTR_tAcc[(None, None, None, subtile_idx + 1)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn_next, next_buf)
acc_vec = tiled_copy_r2s.retile(curr_buf).load()
tRS_rC.store(epilogue_op(acc_vec).to(self.c_dtype))
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)],
)
# Memory ordering for TMA store:
# - `cute.arch.fence_proxy(async_shared)` makes prior SMEM writes visible to async consumers (TMA).
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()
if (subtile_idx + 1) < subtile_cnt:
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()
# Keep warps/roles in lockstep across work items.
self.work_loop_barrier.arrive_and_wait()
# ---------------------------------------------------------------------
# Tail / teardown
# ---------------------------------------------------------------------
if warp_idx == self.tma_warp_id:
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
# Wait for C TMA store complete
c_pipeline.producer_tail()
# Free TMEM
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
def mainloop_s2t_copy_and_partition(self, sSF, tCtSF):
"""Create S2T copy for scale factors.
Uses the tcgen05 S2T (SMEM to TMEM) copy operations for Blackwell.
Following the reference pattern from grouped_blockscaled_gemm.py.
"""
# Filter zeros from both source and destination tensors
# (MMA, MMA_MN, MMA_K, STAGE)
tCsSF_compact = cute.filter_zeros(sSF)
# (MMA, MMA_MN, MMA_K)
tCtSF_compact = cute.filter_zeros(tCtSF)
# Make S2T CopyAtom using Cp4x32x128bOp (32x128b with warpx4 broadcast)
copy_atom_s2t = cute.make_copy_atom(
tcgen05.Cp4x32x128bOp(self.cta_group),
self.sf_dtype,
)
# Create the tiled copy using the tmem tensor
tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
thr_copy_s2t = tiled_copy_s2t.get_slice(0)
# Partition source (SMEM) and get descriptors
# ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
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_
)
# Partition destination (TMEM)
# ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t
def epilog_tmem_copy_and_partition(
self, tidx, tAcc, gC_mnl, epi_tile, use_2cta_instrs
):
"""Partition for TMEM to register copy in epilogue.
Uses sm100_utils.get_tmem_load_op and tcgen05.make_tmem_copy to create
the tiled copy for TMEM to register.
"""
# Make tiledCopy for tensor memory load
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,
)
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE)
tAcc_epi = cute.flat_divide(
tAcc[((None, None), 0, 0, None)],
epi_tile,
)
# (EPI_TILE_M, EPI_TILE_N)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC_mnl_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
# (T2R, T2R_M, T2R_N)
# Double-buffer for TMEM prefetch optimization
tTR_rAcc_0 = cute.make_rmem_tensor(
tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
tTR_rAcc_1 = 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_gC, (tTR_rAcc_0, tTR_rAcc_1)
def epilog_smem_copy_and_partition(self, tiled_copy_t2r, tTR_rC, tidx, sC):
"""Partition for register to SMEM copy in epilogue.
Uses sm100_utils.get_smem_store_op and cute.make_tiled_copy_D.
"""
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)
# (R2S, R2S_M, R2S_N, PIPE_D)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
# (R2S, R2S_M, R2S_N)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(self, tma_atom_c, gC_mnl, epi_tile, sC):
"""Partition for SMEM to GMEM TMA store in epilogue.
Uses cpasync.tma_partition to partition both source (SMEM) and destination (GMEM).
"""
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
# ((ATOM_V, REST_V), EPI_M, EPI_N)
# ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
bSG_sC, bSG_gC = cpasync.tma_partition(
tma_atom_c,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC_for_tma_partition,
)
return tma_atom_c, bSG_sC, bSG_gC
# Global compiled kernel cache
# - key: (m, n, k, l, mma_tiler_mn, cluster_shape_mn, num_ab_stage, num_c_stage)
_compiled_kernel_cache = {}
# Data type constants
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
# Feature flags
USE_PERSISTENT_SCHEDULER = True # Set to False to disable persistent scheduling
# Optional compilation debug flags (off by default).
# Enable with `CUTE_LINEINFO=1` to generate line info in the compiled cubin/ptx.
_ENABLE_CUTE_LINEINFO = os.getenv("CUTE_LINEINFO", "0") == "1"
# Optimal configuration lookup: maps (m, n, k, l) -> (mma_tiler_mn, cluster_shape_mn)
# Enable (1, 2) clustering to use TMA multicast where applicable.
@dataclass(frozen=True)
class KernelConfig:
mma_tiler_mn: Tuple[int, int]
cluster_shape_mn: Tuple[int, int]
num_ab_stage: int = 0 # 0 => auto (use _compute_stages())
num_c_stage: int = 0 # 0 => auto (use _compute_stages())
_OPTIMAL_CONFIG = {
(128, 7168, 16384, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2
),
(128, 4096, 7168, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2
),
(128, 7168, 2048, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2
),
}
# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
(128, 64),
(1, 1),
num_ab_stage=0,
num_c_stage=0,
)
def _env_override_int(name: str) -> Optional[int]:
val = os.getenv(name)
if val is None or val == "":
return None
return int(val)
def _cache_key(
m: int,
n: int,
k: int,
l: int,
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
num_ab_stage: int,
num_c_stage: int,
) -> Tuple[int, int, int, int, Tuple[int, int], Tuple[int, int], int, int]:
return (
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
)
def _ensure_compiled(
m: int,
n: int,
k: int,
l: int,
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
num_ab_stage: int,
num_c_stage: int,
) -> None:
"""Compile a single (shape, tiling) variant if it's missing."""
global _compiled_kernel_cache
key = _cache_key(
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
)
if key in _compiled_kernel_cache:
return
num_ab_stage_override = None if num_ab_stage <= 0 else num_ab_stage
num_c_stage_override = None if num_c_stage <= 0 else num_c_stage
kernel_instance = Sm100BlockScaledGemmKernel(
sf_vec_size=16,
mma_tiler_mn=mma_tiler_mn,
cluster_shape_mn=cluster_shape_mn,
num_ab_stage=num_ab_stage_override,
num_c_stage=num_c_stage_override,
ab_dtype=ab_dtype,
sf_dtype=sf_dtype,
c_dtype=c_dtype,
use_persistent_scheduler=USE_PERSISTENT_SCHEDULER,
)
# Placeholder pointers for compilation
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)
compile_fn = (
cute.compile[cute.GenerateLineInfo(True)]
if _ENABLE_CUTE_LINEINFO
else cute.compile
)
_compiled_kernel_cache[key] = compile_fn(
kernel_instance,
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
(m, n, k, l),
)
def compile_kernel():
"""
Pre-compile ALL required kernel configurations.
This is called by competition code before benchmarking to warm up.
"""
global _compiled_kernel_cache
# Some tests intentionally set `_compiled_kernel_cache = None` to force a cold compile.
# Re-initialize here so both `compile_kernel()` and `custom_kernel()` stay robust.
if _compiled_kernel_cache is None:
_compiled_kernel_cache = {}
# Precompile per-problem variants for the known benchmark shapes.
for shape_key, cfg in _OPTIMAL_CONFIG.items():
m, n, k, l = shape_key
num_ab_stage = max(0, int(cfg.num_ab_stage))
num_c_stage = max(0, int(cfg.num_c_stage))
_ensure_compiled(
m,
n,
k,
l,
cfg.mma_tiler_mn,
cfg.cluster_shape_mn,
num_ab_stage,
num_c_stage,
)
# Conservative fallback so cold shapes still work.
_ensure_compiled(
128,
256,
256,
1,
_DEFAULT_CONFIG.mma_tiler_mn,
_DEFAULT_CONFIG.cluster_shape_mn,
_DEFAULT_CONFIG.num_ab_stage,
_DEFAULT_CONFIG.num_c_stage,
)
# Cache uses default arg trick: mutable default is evaluated once at definition,
# stored in function object, accessed via LOAD_FAST (faster than LOAD_GLOBAL)
_gmem = cute.AddressSpace.gmem
def custom_kernel(data: input_t, _c=[None] * 7) -> output_t:
"""
Execute kernel with pre-compiled configuration from cache.
Cache structure (_c):
- _c[0]: compiled kernel for current key
- _c[1-5]: cached pointer objects (a, b, sfa, sfb, c)
- _c[6]: last cache_key
"""
a, b, _, _, sfa_p, sfb_p, c = data
global _compiled_kernel_cache
if _compiled_kernel_cache is None:
_compiled_kernel_cache = {}
# Get input dimensions
m, n, l = c.shape
k = a.shape[1] << 1 # K is doubled due to FP4 packing
# Get optimal config for this input shape via lookup
shape_key = (m, n, k, l)
cfg = _OPTIMAL_CONFIG.get(shape_key, _DEFAULT_CONFIG)
mma_tiler_mn = cfg.mma_tiler_mn
cluster_shape_mn = cfg.cluster_shape_mn
num_ab_stage = _env_override_int("AB_STAGES")
num_ab_stage = int(cfg.num_ab_stage) if num_ab_stage is None else num_ab_stage
if num_ab_stage < 0:
num_ab_stage = 0
num_c_stage = _env_override_int("C_STAGES")
num_c_stage = int(cfg.num_c_stage) if num_c_stage is None else num_c_stage
if num_c_stage < 0:
num_c_stage = 0
cache_key = _cache_key(
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
)
# Initialize pointer cache on first call (independent of kernel config)
if _c[1] is None:
_c[1] = make_ptr(ab_dtype, 16, _gmem, assumed_align=16)
_c[2] = make_ptr(ab_dtype, 16, _gmem, assumed_align=16)
_c[3] = make_ptr(sf_dtype, 32, _gmem, assumed_align=32)
_c[4] = make_ptr(sf_dtype, 32, _gmem, assumed_align=32)
_c[5] = make_ptr(c_dtype, 16, _gmem, assumed_align=16)
# Swap in compiled kernel when cache key changes
if _c[6] != cache_key or _c[0] is None:
if cache_key not in _compiled_kernel_cache:
_ensure_compiled(
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
)
_c[0] = _compiled_kernel_cache[cache_key]
_c[6] = cache_key
# Update pointer addresses
_c[1]._pointer = a.data_ptr()
_c[1]._c_pointer = None
_c[2]._pointer = b.data_ptr()
_c[2]._c_pointer = None
_c[3]._pointer = sfa_p.data_ptr()
_c[3]._c_pointer = None
_c[4]._pointer = sfb_p.data_ptr()
_c[4]._c_pointer = None
_c[5]._pointer = c.data_ptr()
_c[5]._c_pointer = None
# Launch pre-compiled kernel
_c[0](
_c[1],
_c[2],
_c[3],
_c[4],
_c[5],
(m, n, k, l),
)
return c
scrolls · 1725 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 155546.
⋯ 1475 unchanged lines_OPTIMAL_CONFIG = {(128, 7168, 16384, 1): KernelConfig(- (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2),(128, 4096, 7168, 1): KernelConfig(- (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2),(128, 7168, 2048, 1): KernelConfig(- (128, 64), (1, 2), num_ab_stage=6, num_c_stage=2+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=2),}
Best evidence level for this revision: reported
JSON