submission 154166
rcmalli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2798 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-154166?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:5c317850a9296ad148ccaa25502418d652ab59476a2cc542ac7ab9538b239d58
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 bytessplit-k
workspace_ptr: cute.Pointer, # Workspace for Split-K reductiontcgen05
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.py2798 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 torch
import cutlass
import cutlass.cute as cute
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import llvm
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
# ============================================================================
# dev_v8 developer guidance (comments-only)
# ----------------------------------------------------------------------------
# This file is intended to be modified by the human author. The agent may add
# comments to guide the strm-K / Split-K integration, but should not change
# kernel behavior unless explicitly requested.
#
# Start here:
# - Plan: `dev_v8/PLAN_strmK_SPLITK.md`
# - OpenSpec notes: `openspec/changes/plan-strmk-splitk-dev-v8/`
# - CuTeDSL API list: `dev_v8/CUTEDSL_API_REFERENCE.md`
#
# Primary goal for strm-K/Split-K:
# - Increase parallelism for small-M / large-K shapes (e.g. m=128,n=7168,k=16384)
# where NCU shows a small grid and low waves/SM. strm-K introduces additional
# work items per (m,n,l) tile by splitting the K tiles.
# ============================================================================
@dataclass(frozen=True)
class SchedulerParams:
tiles_m: int
tiles_n: int
k_tile_cnt: int
split_factor: int
@property
def tiles_mn(self) -> int:
return self.tiles_m * self.tiles_n
def total_work_items(self, l: int) -> int:
# DSL-friendly: `split_factor` is the sole control for strm-K/Split-K work expansion.
return self.tiles_mn * l * self.split_factor
@classmethod
def from_shape(
cls,
m: int,
n: int,
k_tile_cnt: int,
tile_m: int,
tile_n: int,
split_factor: 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, split_factor)
@dsl_user_op
def red_global_add_noftz_f16(ptr: cute.Pointer, val, *, loc=None, ip=None) -> None:
"""Atomic reduction add into global memory for FP16 using `red.global.add.noftz.f16`."""
val_bits_i16 = llvm.bitcast(T.i16(), val.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
llvm.inline_asm(
None,
[
ptr.toint(loc=loc, ip=ip).ir_value(loc=loc, ip=ip),
val_bits_i16,
],
"red.global.add.noftz.f16 [$0], $1;",
"l,h",
has_side_effects=True,
asm_dialect=0,
loc=loc,
ip=ip,
)
def _leaf0(x):
"""Extract the first leaf element from a potential XTuple/tuple."""
while cute.depth(x) > 0:
x = cute.get(x, mode=[0])
return x
def _leaf_last(x):
"""Extract the last leaf element from a potential XTuple/tuple."""
while cute.depth(x) > 0:
x = cute.get(x, mode=[cute.rank(x) - 1])
return x
def _leaf_ptr(x):
"""Extract a `cute.Pointer` leaf from a potential XTuple/tuple."""
a = _leaf0(x)
if hasattr(a, "toint"):
return a
b = _leaf_last(x)
if hasattr(b, "toint"):
return b
return a
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,
workspace_ptr: cute.Pointer, # Workspace for Split-K reduction
shape_mnkl: Tuple[int, int, int, int],
split_factor: cutlass.Int32 = 1,
max_active_clusters: cutlass.Constexpr = 132,
epilogue_op: cutlass.Constexpr = lambda x: x,
use_atomic_epilogue: cutlass.Constexpr = False,
):
"""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)
workspace_ptr: Pointer to workspace for Split-K partial sums (M*N*L*(split_factor-1), FP16)
shape_mnkl: Tuple of (M, N, K, L) dimensions
max_active_clusters: Maximum active clusters for persistent scheduling
epilogue_op: Optional epilogue operation
split_factor: Runtime K-split factor (1 disables strm-K/Split-K)
"""
m, n, k, l = shape_mnkl
# strm-K / SPLIT-K HOOK (host-side planning):
# - Define strm-K scheduler params (tiles_m, tiles_n, tiles_k, split_factor, etc.)
# and decide whether to run data-parallel vs split-K/strm-K for this shape.
# - Keep this computation in Python/host path when possible, then pass the params
# into `kernel(...)` as extra arguments.
# - Reference: `dev_v8/PLAN_strmK_SPLITK.md` (Sections 1, 4) and
# `.claude/skills/cute-dsl/cute-dsl-strm-k-scheduling.md`.
# Data types are set during __init__ (self.a_dtype, self.b_dtype, etc.)
# as JIT-compiled code cannot extract element_type from _Pointer objects
# 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)
# PyTorch tensor after permute(1,2,0) has strides (N, 1, M*N)
c_layout = cute.make_layout((m, n, l), stride=(n, 1, m * n))
c_tensor = cute.make_tensor(c_ptr, c_layout)
# Workspace for Split-K reduction: flat buffer of size M*N*L*(split_factor-1)
# Each split partition (except the last) writes its partial sum here
# Layout is flat (1D) - we index by: g_idx + k_part * (M*N*L)
workspace_size = m * n * l * (split_factor - 1) if split_factor > 1 else 1
workspace_layout = cute.make_layout((workspace_size,), stride=(1,))
workspace_tensor = cute.make_tensor(workspace_ptr, workspace_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 (ported from dev_v6)
#
# 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)
# strm-K / Split-K control:
# `split_factor` is a runtime argument so we can tune without recompiling.
# NOTE: Split-K requires a reduction strategy in the epilogue (not implemented yet),
# so callers should keep `split_factor=1` until that work lands.
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,
split_factor=split_factor,
)
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
# Launch kernel (two isolated entrypoints)
if cutlass.const_expr(use_atomic_epilogue):
self.kernel_split(
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,
workspace_tensor, # Workspace for Split-K reduction
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,
split_factor=split_factor,
).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,
)
else:
self.kernel_nosplit(
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.jit
def call_tma(
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,
):
# Dedicated TMA epilogue launcher (no Split-K).
# Note: workspace_ptr is not used when split_factor=1, but we need to pass
# a valid pointer to satisfy __call__ signature. Use c_ptr as dummy.
return self(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
c_ptr, # workspace_ptr (unused, split_factor=1)
shape_mnkl,
cutlass.Int32(1),
max_active_clusters=max_active_clusters,
epilogue_op=epilogue_op,
)
@cute.jit
def call_atomic(
self,
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
workspace_ptr: cute.Pointer, # Workspace for Split-K reduction
shape_mnkl: Tuple[int, int, int, int],
split_factor: cutlass.Constexpr,
max_active_clusters: cutlass.Constexpr = 132,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
# Dedicated workspace reduction launcher (Split-K / strm-K partial sums).
return self(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
workspace_ptr,
shape_mnkl,
split_factor,
max_active_clusters=max_active_clusters,
epilogue_op=epilogue_op,
)
@staticmethod
@cute.jit
def decode_work_idx(
work_idx,
tiles_m,
tiles_n,
tiles_mn,
k_tile_cnt,
split_factor: cutlass.Constexpr,
):
work_m_idx = work_idx % tiles_m
work_n_idx = (work_idx // tiles_m) % tiles_n
work_l_idx = work_idx // tiles_mn
# strm-K split along K tiles (even partition; empty ranges allowed if split_factor > k_tile_cnt)
k_part = (work_idx // tiles_mn) % split_factor
k_tile_start = (k_part * k_tile_cnt) // split_factor
k_tile_end = ((k_part + 1) * k_tile_cnt) // split_factor
# Return k_part for workspace reduction (split index: 0..split_factor-1)
return work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, k_part
@cute.kernel
def kernel_nosplit(
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)
@cute.kernel
def kernel_split(
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,
mWorkspace_mnl: cute.Tensor, # Workspace tensor: (M, N, L, split_factor-1)
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,
split_factor: cutlass.Constexpr = 1,
):
"""
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.
"""
use_atomic_epilogue = True
# 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
)
if cutlass.const_expr(not use_atomic_epilogue):
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
# strm-K / SPLIT-K HOOK (device-side decode):
# Replace the simple (m,n,l) decode below with a strm-K decode that returns:
# (work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, is_splitk)
#
# Minimal first step (no behavior change):
# - Implement decode functions but keep `k_tile_start=0`, `k_tile_end=k_tile_cnt`,
# `is_splitk=False` so correctness stays identical.
#
# Full strm-K step:
# - Use `k_tile_start/k_tile_end` to restrict the K-tile loop for both TMA and MMA.
# - Ensure K coverage invariant holds (see plan + `.claude/.../cute-dsl-kernel-testing.md`).
work_m_idx, work_n_idx, work_l_idx, k_tile_start, k_tile_end, k_part = (
self.decode_work_idx(
work_idx,
tiles_m,
tiles_n,
tiles_mn,
k_tile_cnt,
split_factor,
)
)
k_tile_cnt_local = k_tile_end - k_tile_start
# Workspace reduction: last split writes to output, others to workspace
is_last_split = k_part == split_factor - 1
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_local:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
# strm-K / SPLIT-K HOOK:
# For split-K, change the loop to iterate only over the assigned K tile range:
# for _ in range(k_tile_cnt_local):
# k_tile = k_tile_start + ab_producer_state.count
# and index gmem tiles using `k_tile` (not raw `count`).
for k_iter in cutlass.range(0, k_tile_cnt_local, 1, unroll=1): # noqa: B007
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
k_tile = k_tile_start + 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_local:
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_local 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)
# strm-K / SPLIT-K HOOK:
# For split-K, iterate only the assigned K-tile subset.
# IMPORTANT: You still want `Field.ACCUMULATE=False` for the first kblock of the
# first K tile in THIS CTA's range, then `True` for subsequent kblocks/K-tiles.
for k_iter in range(k_tile_cnt_local): # 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_local:
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:
# strm-K / SPLIT-K HOOK:
# For split-K tiles, this epilogue cannot directly TMA-store to `C` because multiple
# CTAs contribute partial sums to the same output tile.
#
# Typical workspace-free strategy: convert fragments to FP16/FP32 and use global atomics
# to accumulate into `C`. This requires guaranteeing `C` is initialized to zero for
# split-K tiles (or implementing a safe init protocol).
#
# Reference:
# - `dev_v8/PLAN_strmK_SPLITK.md` (Output semantics)
# - `.claude/skills/cute-dsl/cute-dsl-nvfp4-gemv-optimization.md` (nvvm.atomicrmw FADD pattern)
# - `dev_v8/CUTEDSL_API_REFERENCE.md` (atomics gap section)
if cutlass.const_expr(not use_atomic_epilogue):
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))
# Slice to the current output tile so coordinate-based addressing works for Split-K atomics.
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))
if cutlass.const_expr(not use_atomic_epilogue):
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 cutlass.const_expr(not use_atomic_epilogue):
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)
if cutlass.const_expr(use_atomic_epilogue):
# Split-K epilogue with in-kernel workspace reduction (no atomics).
# - Splits 0 to (N-2): write partials to workspace[k_part]
# - Split (N-1): read all workspace partials, sum with own partial, write to C
# This eliminates the need for a separate reduction kernel.
tTR_gC_mn = tTR_gC[(None, None, None, subtile_idx)]
elem_cnt = cute.size(curr_buf.shape)
output_size = cute.size(mC_raw_mnl.layout)
# Element-wise processing (no atomics) - hardware will coalesce
for elem_i in cutlass.range(elem_cnt, unroll_full=True):
crd = cute.idx2crd(elem_i, curr_buf.shape)
val_acc = curr_buf[crd]
val_c = epilogue_op(val_acc).to(self.c_dtype)
g_crd = tTR_gC_mn[crd]
g_idx = cute.crd2idx(g_crd, mC_raw_mnl.layout)
if is_last_split:
# Last split: read all workspace partials, sum, write final result
# Accumulate partials from all previous splits
total = val_c
for p in range(split_factor - 1):
workspace_val = mWorkspace_mnl[
g_idx + p * output_size
]
total = total + workspace_val
mC_raw_mnl[g_idx] = total
else:
# Other splits: write to workspace slice
# Workspace offset: k_part * (M*N*L)
mWorkspace_mnl[g_idx + k_part * output_size] = val_c
else:
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:
if cutlass.const_expr(not use_atomic_epilogue):
# 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, split_factor)
_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]
split_factor: int = 1
num_ab_stage: int = 0 # 0 => auto (use _compute_stages())
num_c_stage: int = 0 # 0 => auto (use _compute_stages())
_OPTIMAL_CONFIG = {
# You manage all tuning here per problem type:
# KernelConfig(mma_tiler_mn=(tile_m, tile_n), cluster_shape_mn=(cluster_m, cluster_n),
# split_factor=..., num_ab_stage=...)
(128, 7168, 16384, 1): KernelConfig(
(128, 128), (1, 2), split_factor=1, num_ab_stage=0
),
(128, 4096, 7168, 1): KernelConfig(
(128, 64), (1, 2), split_factor=1, num_ab_stage=0
),
(128, 7168, 2048, 1): KernelConfig(
(128, 64), (1, 2), split_factor=1, num_ab_stage=0
),
}
# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
(128, 64),
(1, 1),
split_factor=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],
split_factor: int,
num_ab_stage: int,
num_c_stage: int,
) -> Tuple[int, int, int, int, Tuple[int, int], Tuple[int, int], int, int, int]:
return (
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
split_factor,
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],
split_factor: int,
num_ab_stage: int,
num_c_stage: int,
) -> None:
"""Compile a single (shape, tiling, split_factor) variant if it's missing."""
global _compiled_kernel_cache
key = _cache_key(
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
split_factor,
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)
workspace_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
compile_fn = (
cute.compile[cute.GenerateLineInfo(True)]
if _ENABLE_CUTE_LINEINFO
else cute.compile
)
if split_factor > 1:
_compiled_kernel_cache[key] = compile_fn(
kernel_instance.call_atomic,
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
workspace_ptr,
(m, n, k, l),
split_factor,
)
else:
_compiled_kernel_cache[key] = compile_fn(
kernel_instance.call_tma,
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
split_factor = max(1, int(cfg.split_factor))
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,
split_factor,
num_ab_stage,
num_c_stage,
)
# Conservative fallback (split_factor=1) so cold shapes still work.
_ensure_compiled(
128,
256,
256,
1,
_DEFAULT_CONFIG.mma_tiler_mn,
_DEFAULT_CONFIG.cluster_shape_mn,
_DEFAULT_CONFIG.split_factor,
_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] * 9) -> 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
- _c[7]: workspace pointer
- _c[8]: workspace tensor (reused across calls)
"""
a, b, _, _, sfa_p, sfb_p, c = data
global _compiled_kernel_cache
if _compiled_kernel_cache is None:
_compiled_kernel_cache = {}
# -------------------------------------------------------------------------
# Input / output tensor shapes from the task harness:
#
# Inputs:
# - `a`: (M, K_packed, L) nvfp4(e2m1), K-major. NOTE: `K_packed = K//2` (FP4 packed).
# - `b`: (N, K_packed, L) nvfp4(e2m1), K-major. NOTE: `K_packed = K//2` (FP4 packed).
# - `sfa`: (M, K//16, L) fp8(e4m3fnuz) reference layout (NOT used by this kernel).
# - `sfb`: (N, K//16, L) fp8(e4m3fnuz) reference layout (NOT used by this kernel).
# - `sfa_p` (`sfa_permuted`): CuTe kernel scale-factor layout for A, typically
# shaped like (32, 4, rest_m, 4, rest_k, L) (exact factorization depends on tiling).
# - `sfb_p` (`sfb_permuted`): CuTe kernel scale-factor layout for B, typically
# shaped like (32, 4, rest_n, 4, rest_k, L).
#
# Output:
# - `c`: (M, N, L) fp16 output tile accumulator / destination.
# -------------------------------------------------------------------------
# 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
# strm-K / Split-K tuning parameter.
# NOTE: `split_factor` is compile-time-specialized (part of the cache key) for max performance.
# Changing it at runtime selects/compiles a different cached kernel variant.
split_factor = _env_override_int("SPLIT_FACTOR")
split_factor = int(cfg.split_factor) if split_factor is None else split_factor
if split_factor < 1:
split_factor = 1
# Note: With workspace reduction, C does NOT need to be zero-initialized.
# The last split writes directly to C (non-atomic), and workspace partials are summed.
use_atomic_epilogue = split_factor > 1
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,
split_factor,
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)
_c[7] = make_ptr(c_dtype, 16, _gmem, assumed_align=16) # workspace pointer
# 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,
split_factor,
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
if use_atomic_epilogue:
# Workspace reduction path: allocate workspace, run GEMM kernel
# The kernel performs in-kernel reduction: last split reads workspace partials,
# sums them with its own partial, and writes the final result to C.
# No separate reduction kernel needed!
workspace_size = m * n * l * (split_factor - 1)
# Allocate or reuse workspace tensor
if _c[8] is None or _c[8].numel() < workspace_size:
_c[8] = torch.zeros(workspace_size, dtype=torch.float16, device=c.device)
workspace = _c[8][:workspace_size]
# Update workspace pointer
_c[7]._pointer = workspace.data_ptr()
_c[7]._c_pointer = None
# Run GEMM kernel - in-kernel reduction handles everything
_c[0](_c[1], _c[2], _c[3], _c[4], _c[5], _c[7], (m, n, k, l))
else:
_c[0](
_c[1],
_c[2],
_c[3],
_c[4],
_c[5],
(m, n, k, l),
)
return c
scrolls · 2798 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 126536.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON