submission 168494
rcmalli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1697 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-168494?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:8f31d8dbf2391801aa1ce1ef9c98ed9f91fd3f755781f4eb458d57a710cd3348
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
- Independent TMA/MMA/Epilogue work loops (no global sync barrier)mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(persistent-kernel
This implements a persistent warp-specialized kernel for B200 (SM100) targetingshared-memory
def make_smem_layout_sfa_custom(tcgen05
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONEwarp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
submission.py1697 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.
V10: Clean rewrite with independent work loops per warp role
- Independent TMA/MMA/Epilogue work loops (no global sync barrier)
- StaticPersistentTileScheduler per warp role
- TMEM allocation only for MMA and Epilogue warps
- Proper barrier scoping
"""
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
# =============================================================================
# Custom SMEM Layout Functions for Scale Factors
# These are modified versions of blockscaled_utils functions that accept
# mma_inst_tile_k as a parameter instead of hardcoding it to 4.
# =============================================================================
@dataclass(frozen=True)
class BlockScaledBasicChunk:
"""
The basic scale factor atom layout decided by tcgen05 BlockScaled MMA Ops.
Copied from cutlass.utils.blockscaled_layout to avoid import issues.
"""
sf_vec_size: int
@property
def layout(self) -> cute.Layout:
# K-major layout: (AtomMN, AtomK)
atom_shape = ((32, 4), (self.sf_vec_size, 4))
atom_stride = ((16, 4), (0, 1))
return cute.make_layout(atom_shape, stride=atom_stride)
def make_smem_layout_sfa_custom(
tiled_mma: cute.TiledMma,
mma_tiler_mnk: Tuple[int, int, int],
sf_vec_size: int,
num_stages: int,
mma_inst_tile_k: int = 4, # NEW: Configurable parameter
) -> cute.Layout:
"""
Custom version of blockscaled_utils.make_smem_layout_sfa that accepts
mma_inst_tile_k as a parameter instead of hardcoding it to 4.
"""
# (CTA_Tile_Shape_M, MMA_Tile_Shape_K)
sfa_tile_shape = (
mma_tiler_mnk[0] // cute.size(tiled_mma.thr_id.shape),
mma_tiler_mnk[2],
)
# ((Atom_M, Rest_M),(Atom_K, Rest_K))
smem_layout = cute.tile_to_shape(
BlockScaledBasicChunk(sf_vec_size).layout,
sfa_tile_shape,
(2, 1),
)
# Use the passed mma_inst_tile_k instead of hardcoded 4
# (CTA_Tile_Shape_M, MMA_Inst_Shape_K)
sfa_tile_shape = cute.shape_div(sfa_tile_shape, (1, mma_inst_tile_k))
# ((Atom_Inst_M, Atom_Inst_K), MMA_M, MMA_K))
smem_layout = cute.tiled_divide(smem_layout, sfa_tile_shape)
atom_m = 128
tiler_inst = ((atom_m, sf_vec_size),)
# (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)
smem_layout = cute.logical_divide(smem_layout, tiler_inst)
# (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE)
sfa_smem_layout_staged = cute.append(
smem_layout,
cute.make_layout(
num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
),
)
return sfa_smem_layout_staged
def make_smem_layout_sfb_custom(
tiled_mma: cute.TiledMma,
mma_tiler_mnk: Tuple[int, int, int],
sf_vec_size: int,
num_stages: int,
mma_inst_tile_k: int = 4, # NEW: Configurable parameter
) -> cute.Layout:
"""
Custom version of blockscaled_utils.make_smem_layout_sfb that accepts
mma_inst_tile_k as a parameter instead of hardcoding it to 4.
"""
# (Round_Up(CTA_Tile_Shape_N, 128), MMA_Tile_Shape_K)
sfb_tile_shape = (
cute.round_up(mma_tiler_mnk[1], 128),
mma_tiler_mnk[2],
)
# ((Atom_N, Rest_N),(Atom_K, Rest_K))
smem_layout = cute.tile_to_shape(
BlockScaledBasicChunk(sf_vec_size).layout,
sfb_tile_shape,
(2, 1),
)
# Use the passed mma_inst_tile_k instead of hardcoded 4
# (CTA_Tile_Shape_N, MMA_Inst_Shape_K)
sfb_tile_shape = cute.shape_div(sfb_tile_shape, (1, mma_inst_tile_k))
# ((Atom_Inst_N, Atom_Inst_K), MMA_N, MMA_K)
smem_layout = cute.tiled_divide(smem_layout, sfb_tile_shape)
atom_n = 128
tiler_inst = ((atom_n, sf_vec_size),)
# (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)
smem_layout = cute.logical_divide(smem_layout, tiler_inst)
# (((Atom_Inst_N, Rest_N),(Atom_Inst_K, Rest_K)), MMA_N, MMA_K, STAGE)
sfb_smem_layout_staged = cute.append(
smem_layout,
cute.make_layout(
num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
),
)
return sfb_smem_layout_staged
class Sm100BlockScaledGemmKernel:
"""
Block-scaled FP4 GEMM kernel for SM100 (Blackwell).
V10 Architecture:
- Independent work loops per warp role (TMA, MMA, Epilogue)
- No global synchronization barrier between work items
- StaticPersistentTileScheduler per warp role
- TMEM allocation only for MMA and Epilogue warps
"""
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,
mma_inst_tile_k: int = 4, # NEW: Configurable K-tile multiplier
ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN,
sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN,
c_dtype: Type[cutlass.Numeric] = cutlass.Float16,
):
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.mma_inst_tile_k = mma_inst_tile_k # Store for use in layout functions
self.num_ab_stage_override = num_ab_stage
self.num_c_stage_override = num_c_stage
# Store data types as instance variables
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
# Epilogue sync barrier - only epilogue warps participate
self.epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=1,
num_threads=32 * len(self.epilog_warp_id),
)
# TMEM allocation barrier - CRITICAL: Only MMA + Epilogue warps (excludes TMA warp)
self.tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=2,
num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
)
# SM100 has 228KB shared memory
self.smem_capacity = 228 * 1024
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])
# Use configurable mma_inst_tile_k from __init__
self.mma_tiler = (
self.mma_inst_shape_mn[0],
self.mma_inst_shape_mn[1],
mma_inst_shape_k * self.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 * self.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(30, 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,
)
# Use custom SMEM layout functions that accept mma_inst_tile_k
self.sfa_smem_layout_staged = make_smem_layout_sfa_custom(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
self.mma_inst_tile_k,
)
self.sfb_smem_layout_staged = make_smem_layout_sfb_custom(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
self.mma_inst_tile_k,
)
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
) * self.mma_inst_tile_k
self.num_sfb_tmem_cols = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
) * self.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)
# Use custom layout functions with configurable mma_inst_tile_k
sfa_smem_layout = make_smem_layout_sfa_custom(
tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k
)
sfb_smem_layout = make_smem_layout_sfb_custom(
tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k
)
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
# Use num_acc_stage=1 for overlapping accumulator
num_acc_stage = 1
# Compute C stages (2 for double buffering)
num_c_stage = 2
# Fixed overhead for barriers
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(30, available // ab_sf_bytes_per_stage))
return num_acc_stage, num_ab_stage, num_c_stage
@cute.jit
def __call__(
self,
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
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."""
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
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
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,
)
# Compute grid using PersistentTileScheduler
c_shape = cute.slice_(self.cta_tile_shape_mnk, (None, None, 0))
gc = cute.zipped_divide(c_tensor, tiler=c_shape)
num_ctas_mnl = gc[(0, (None, None, None))].shape
cluster_shape_mnk = (*self.cluster_shape_mn, 1)
tile_sched_params = utils.PersistentTileSchedulerParams(
num_ctas_mnl, cluster_shape_mnk
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
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,
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,
tile_sched_params,
epilogue_op,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 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,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
tile_sched_params: utils.PersistentTileSchedulerParams,
epilogue_op: cutlass.Constexpr,
):
"""
GPU device kernel for block-scaled GEMM with independent work loops per warp role.
"""
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# TMA descriptor prefetch
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
bidx, bidy, bidz = cute.arch.block_idx()
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
cta_rank_in_cluster
)
tidx, _, _ = cute.arch.thread_idx()
# Shared memory allocation
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
# Initialize mainloop ab_pipeline
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,
)
# Initialize acc_pipeline
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
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
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
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(gA_mkl, mode=[3])
# 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 for A
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)
# =================================================================
# INDEPENDENT WORK LOOPS PER WARP ROLE
# No global synchronization barrier between work items
# =================================================================
# -----------------------------------------------------------------
# TMA LOAD WARP (warp 5) - INDEPENDENT WORK LOOP
# -----------------------------------------------------------------
if warp_idx == self.tma_warp_id:
# Create tile scheduler for TMA warp
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile = tile_sched.initial_work_tile_info()
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
while work_tile.is_valid_tile:
# Get tile coord from tile scheduler
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
)
# Slice to per mma tile index
tAgA_slice = tAgA[
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB[
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA[
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]
# Reset and peek for first k tile
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
)
# TMA load loop
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
# Wait for AB buffer empty
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA loads
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
)
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
)
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
)
# Advance and peek for next k tile
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
)
# Advance to next tile
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
# TMA warp tail
ab_pipeline.producer_tail(ab_producer_state)
# -----------------------------------------------------------------
# MMA WARP (warp 4) - INDEPENDENT WORK LOOP
# -----------------------------------------------------------------
if warp_idx == self.mma_warp_id:
# Wait for TMEM allocation
tmem.wait_for_alloc()
# Retrieve TMEM pointers
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
# Make SFA tmem tensor
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)
# Make SFB tmem tensor
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)
# S2T copy partitions for scale factors
(
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)
# Create tile scheduler for MMA warp
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
while work_tile.is_valid_tile:
# Get tile coord from tile scheduler
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
)
# Get accumulator stage index
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)]
# Peek for first k tile
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
)
# Wait for accumulator buffer empty
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
# Handle SFB offset for different tile shapes
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)
# Reset accumulate flag
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
# MMA mainloop
for k_tile in range(k_tile_cnt):
if is_leader_cta:
# Wait for AB buffer full
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
# Copy SFA/SFB from smem to tmem
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,
)
# MMA over k blocks
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (
None,
None,
kblock_idx,
ab_consumer_state.index,
)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(
tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator
)
tiled_mma.set(
tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator
)
cute.gemm(
tiled_mma,
tCtAcc,
tCrA[kblock_coord],
tCrB[kblock_coord],
tCtAcc,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
# Release AB buffer
ab_pipeline.consumer_release(ab_consumer_state)
# Peek for next k tile
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
)
# Commit accumulator
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
# Advance to next tile
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
# MMA warp tail
acc_pipeline.producer_tail(acc_producer_state)
# -----------------------------------------------------------------
# EPILOGUE WARPS (warps 0-3) - INDEPENDENT WORK LOOP
# -----------------------------------------------------------------
if warp_idx < self.mma_warp_id:
# Allocate TMEM
tmem.allocate(self.num_tmem_alloc_cols)
# Wait for TMEM allocation
tmem.wait_for_alloc()
# Retrieve TMEM pointer
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
# Partition for epilogue
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_rAcc,
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
(
tma_atom_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(tma_atom_c, tCgC, epi_tile, sC)
# Create tile scheduler for Epilogue warps
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile = tile_sched.initial_work_tile_info()
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
# C TMA store pipeline
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage,
producer_group=c_producer_group,
)
while work_tile.is_valid_tile:
# Get tile coord from tile scheduler
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
)
# Slice output tensor
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
# Get accumulator stage index
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
reverse_subtile = (
cutlass.Boolean(True)
if acc_stage_index == 0
else cutlass.Boolean(False)
)
else:
acc_stage_index = acc_consumer_state.index
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
# Wait for accumulator buffer full
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
# Store accumulator to global memory in subtiles
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in cutlass.range(subtile_cnt):
real_subtile_idx = subtile_idx
if cutlass.const_expr(self.overlapping_accum):
if reverse_subtile:
real_subtile_idx = (
self.cta_tile_shape_mnk[1] // self.epi_tile_n
- 1
- subtile_idx
)
# Load accumulator from TMEM to registers
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
# Early release for overlapping accum
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
# Convert to C type
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
tRS_rC.store(acc_vec)
# Store C to shared memory
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
# Fence and barrier
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
# TMA store C
if warp_idx == self.epilog_warp_id[0]:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
# Release accumulator buffer (non-overlapping case)
if cutlass.const_expr(not self.overlapping_accum):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
# Advance to next tile
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
# Epilogue tail - dealloc TMEM and wait for C stores
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
c_pipeline.producer_tail()
def mainloop_s2t_copy_and_partition(self, sSF, tCtSF):
"""Create S2T copy for scale factors."""
tCsSF_compact = cute.filter_zeros(sSF)
tCtSF_compact = cute.filter_zeros(tCtSF)
copy_atom_s2t = cute.make_copy_atom(
tcgen05.Cp4x32x128bOp(self.cta_group),
self.sf_dtype,
)
tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
thr_copy_s2t = tiled_copy_s2t.get_slice(0)
tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t, tCsSF_compact_s2t_
)
tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t
def epilog_tmem_copy_and_partition(
self, tidx, tAcc, gC_mnl, epi_tile, use_2cta_instrs
):
"""Partition for TMEM to register copy in epilogue."""
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
tAcc_epi = cute.flat_divide(
tAcc[((None, None), 0, 0, None)],
epi_tile,
)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
gC_mnl_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
tTR_rAcc = cute.make_rmem_tensor(
tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
return tiled_copy_t2r, tTR_tAcc, tTR_rAcc
def epilog_smem_copy_and_partition(self, tiled_copy_t2r, tTR_rC, tidx, sC):
"""Partition for register to SMEM copy in epilogue."""
copy_atom_r2s = sm100_utils.get_smem_store_op(
self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
)
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(self, tma_atom_c, gC_mnl, epi_tile, sC):
"""Partition for SMEM to GMEM TMA store in epilogue."""
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)
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
# =============================================================================
# Kernel Compilation and Dispatch
# =============================================================================
_compiled_kernel_cache = {}
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
@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())
mma_inst_tile_k: int = 4 # K-tile multiplier: 1, 2, or 4
_OPTIMAL_CONFIG = {
(128, 7168, 16384, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
),
(128, 4096, 7168, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
),
(128, 7168, 2048, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=2, mma_inst_tile_k=4
),
}
# Default fallback configuration for unseen shapes.
_DEFAULT_CONFIG = KernelConfig(
(128, 64),
(1, 1),
num_ab_stage=7,
num_c_stage=2,
mma_inst_tile_k=4,
)
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,
mma_inst_tile_k: int,
) -> Tuple:
return (
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
mma_inst_tile_k,
)
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,
mma_inst_tile_k: 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,
mma_inst_tile_k,
)
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,
mma_inst_tile_k=mma_inst_tile_k,
ab_dtype=ab_dtype,
sf_dtype=sf_dtype,
c_dtype=c_dtype,
)
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)
_compiled_kernel_cache[key] = cute.compile(
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."""
global _compiled_kernel_cache
if _compiled_kernel_cache is None:
_compiled_kernel_cache = {}
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))
mma_inst_tile_k = int(cfg.mma_inst_tile_k)
_ensure_compiled(
m,
n,
k,
l,
cfg.mma_tiler_mn,
cfg.cluster_shape_mn,
num_ab_stage,
num_c_stage,
mma_inst_tile_k,
)
_gmem = cute.AddressSpace.gmem
def custom_kernel(data: input_t, _c=[None] * 7) -> output_t:
"""Execute kernel with pre-compiled configuration from cache."""
a, b, _, _, sfa_p, sfb_p, c = data
global _compiled_kernel_cache
if _compiled_kernel_cache is None:
_compiled_kernel_cache = {}
m, n, l = c.shape
k = a.shape[1] << 1
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
mma_inst_tile_k = int(cfg.mma_inst_tile_k)
cache_key = _cache_key(
m,
n,
k,
l,
mma_tiler_mn,
cluster_shape_mn,
num_ab_stage,
num_c_stage,
mma_inst_tile_k,
)
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)
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,
mma_inst_tile_k,
)
_c[0] = _compiled_kernel_cache[cache_key]
_c[6] = cache_key
_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
_c[0](
_c[1],
_c[2],
_c[3],
_c[4],
_c[5],
(m, n, k, l),
)
return c
scrolls · 1697 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 155565.
⋯ 3 unchanged linesThis implements a persistent warp-specialized kernel for B200 (SM100) targetingthe 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+ V10: Clean rewrite with independent work loops per warp role+ - Independent TMA/MMA/Epilogue work loops (no global sync barrier)+ - StaticPersistentTileScheduler per warp role+ - TMEM allocation only for MMA and Epilogue warps+ - Proper barrier scoping"""from dataclasses import dataclass⋯ 14 unchanged linesfrom task import input_t, output_t+ # =============================================================================+ # Custom SMEM Layout Functions for Scale Factors+ # These are modified versions of blockscaled_utils functions that accept+ # mma_inst_tile_k as a parameter instead of hardcoding it to 4.+ # =============================================================================++@dataclass(frozen=True)- class SchedulerParams:- tiles_m: int- tiles_n: int- k_tile_cnt: int+ class BlockScaledBasicChunk:+ """+ The basic scale factor atom layout decided by tcgen05 BlockScaled MMA Ops.+ Copied from cutlass.utils.blockscaled_layout to avoid import issues.+ """+ sf_vec_size: int+@property- def tiles_mn(self) -> int:- return self.tiles_m * self.tiles_n+ def layout(self) -> cute.Layout:+ # K-major layout: (AtomMN, AtomK)+ atom_shape = ((32, 4), (self.sf_vec_size, 4))+ atom_stride = ((16, 4), (0, 1))+ return cute.make_layout(atom_shape, stride=atom_stride)- 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)+ def make_smem_layout_sfa_custom(+ tiled_mma: cute.TiledMma,+ mma_tiler_mnk: Tuple[int, int, int],+ sf_vec_size: int,+ num_stages: int,+ mma_inst_tile_k: int = 4, # NEW: Configurable parameter+ ) -> cute.Layout:+ """+ Custom version of blockscaled_utils.make_smem_layout_sfa that accepts+ mma_inst_tile_k as a parameter instead of hardcoding it to 4.+ """+ # (CTA_Tile_Shape_M, MMA_Tile_Shape_K)+ sfa_tile_shape = (+ mma_tiler_mnk[0] // cute.size(tiled_mma.thr_id.shape),+ mma_tiler_mnk[2],+ )+ # ((Atom_M, Rest_M),(Atom_K, Rest_K))+ smem_layout = cute.tile_to_shape(+ BlockScaledBasicChunk(sf_vec_size).layout,+ sfa_tile_shape,+ (2, 1),+ )+ # Use the passed mma_inst_tile_k instead of hardcoded 4+ # (CTA_Tile_Shape_M, MMA_Inst_Shape_K)+ sfa_tile_shape = cute.shape_div(sfa_tile_shape, (1, mma_inst_tile_k))+ # ((Atom_Inst_M, Atom_Inst_K), MMA_M, MMA_K))+ smem_layout = cute.tiled_divide(smem_layout, sfa_tile_shape)++ atom_m = 128+ tiler_inst = ((atom_m, sf_vec_size),)+ # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)+ smem_layout = cute.logical_divide(smem_layout, tiler_inst)++ # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE)+ sfa_smem_layout_staged = cute.append(+ smem_layout,+ cute.make_layout(+ num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))+ ),+ )++ return sfa_smem_layout_staged+++ def make_smem_layout_sfb_custom(+ tiled_mma: cute.TiledMma,+ mma_tiler_mnk: Tuple[int, int, int],+ sf_vec_size: int,+ num_stages: int,+ mma_inst_tile_k: int = 4, # NEW: Configurable parameter+ ) -> cute.Layout:+ """+ Custom version of blockscaled_utils.make_smem_layout_sfb that accepts+ mma_inst_tile_k as a parameter instead of hardcoding it to 4.+ """+ # (Round_Up(CTA_Tile_Shape_N, 128), MMA_Tile_Shape_K)+ sfb_tile_shape = (+ cute.round_up(mma_tiler_mnk[1], 128),+ mma_tiler_mnk[2],+ )++ # ((Atom_N, Rest_N),(Atom_K, Rest_K))+ smem_layout = cute.tile_to_shape(+ BlockScaledBasicChunk(sf_vec_size).layout,+ sfb_tile_shape,+ (2, 1),+ )++ # Use the passed mma_inst_tile_k instead of hardcoded 4+ # (CTA_Tile_Shape_N, MMA_Inst_Shape_K)+ sfb_tile_shape = cute.shape_div(sfb_tile_shape, (1, mma_inst_tile_k))+ # ((Atom_Inst_N, Atom_Inst_K), MMA_N, MMA_K)+ smem_layout = cute.tiled_divide(smem_layout, sfb_tile_shape)++ atom_n = 128+ tiler_inst = ((atom_n, sf_vec_size),)+ # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K)+ smem_layout = cute.logical_divide(smem_layout, tiler_inst)++ # (((Atom_Inst_N, Rest_N),(Atom_Inst_K, Rest_K)), MMA_N, MMA_K, STAGE)+ sfb_smem_layout_staged = cute.append(+ smem_layout,+ cute.make_layout(+ num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))+ ),+ )++ return sfb_smem_layout_staged++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)+ V10 Architecture:+ - Independent work loops per warp role (TMA, MMA, Epilogue)+ - No global synchronization barrier between work items+ - StaticPersistentTileScheduler per warp role+ - TMEM allocation only for MMA and Epilogue warps"""def __init__(⋯ 3 unchanged linescluster_shape_mn: Tuple[int, int] = (1, 2),num_ab_stage: Optional[int] = None,num_c_stage: Optional[int] = None,+ mma_inst_tile_k: int = 4, # NEW: Configurable K-tile multiplierab_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.Float32self.sf_vec_size = sf_vec_sizeself.use_2cta_instrs = mma_tiler_mn[0] == 256self.cluster_shape_mn = cluster_shape_mnself.mma_tiler = (*mma_tiler_mn, 1) # K dimension set later- self.use_persistent_scheduler = use_persistent_scheduler+ self.mma_inst_tile_k = mma_inst_tile_k # Store for use in layout functionsself.num_ab_stage_override = num_ab_stageself.num_c_stage_override = num_c_stage- # Store data types as instance variables (needed for JIT compilation)+ # Store data types as instance variablesself.a_dtype = ab_dtypeself.b_dtype = ab_dtypeself.sf_dtype = sf_dtype⋯ 14 unchanged lines)# Named barriers+ # Epilogue sync barrier - only epilogue warps participateself.epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1,num_threads=32 * len(self.epilog_warp_id),)+ # TMEM allocation barrier - CRITICAL: Only MMA + Epilogue warps (excludes TMA warp)self.tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=2,- num_threads=self.threads_per_cta,+ num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),)- # 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 has 228KB shared memory+ self.smem_capacity = 228 * 1024SM100_TMEM_CAPACITY_COLUMNS = 512self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS⋯ 29 unchanged lines# Compute MMA/cluster/tile shapesmma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])- mma_inst_tile_k = 4+ # Use configurable mma_inst_tile_k from __init__self.mma_tiler = (self.mma_inst_shape_mn[0],self.mma_inst_shape_mn[1],- mma_inst_shape_k * mma_inst_tile_k,+ mma_inst_shape_k * self.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,+ mma_inst_shape_k * self.mma_inst_tile_k,)self.cta_tile_shape_mnk = (self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),⋯ 48 unchanged linesself.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)))+ self.num_ab_stage = max(2, min(30, 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)))⋯ 10 unchanged linesself.b_dtype,self.num_ab_stage,)- self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(+ # Use custom SMEM layout functions that accept mma_inst_tile_k+ self.sfa_smem_layout_staged = make_smem_layout_sfa_custom(tiled_mma,self.mma_tiler,self.sf_vec_size,self.num_ab_stage,+ self.mma_inst_tile_k,)- self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(+ self.sfb_smem_layout_staged = make_smem_layout_sfb_custom(tiled_mma,self.mma_tiler,self.sf_vec_size,self.num_ab_stage,+ self.mma_inst_tile_k,)self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(self.c_dtype,⋯ 9 unchanged linessf_atom_mn = 32self.num_sfa_tmem_cols = (self.cta_tile_shape_mnk[0] // sf_atom_mn- ) * mma_inst_tile_k+ ) * self.mma_inst_tile_kself.num_sfb_tmem_cols = (self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn- ) * mma_inst_tile_k+ ) * self.mma_inst_tile_kself.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_colsself.num_accumulator_tmem_cols = (self.cta_tile_shape_mnk[1] * self.num_acc_stage⋯ 22 unchanged lines# Compute SMEM footprintsa_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+ # Use custom layout functions with configurable mma_inst_tile_k+ sfa_smem_layout = make_smem_layout_sfa_custom(+ tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k)- sfb_smem_layout = blockscaled_utils.make_smem_layout_sfb(- tiled_mma, mma_tiler, sf_vec_size, 1+ sfb_smem_layout = make_smem_layout_sfb_custom(+ tiled_mma, mma_tiler, sf_vec_size, 1, self.mma_inst_tile_k)c_smem_layout = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)⋯ 5 unchanged linesab_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+ # Use num_acc_stage=1 for overlapping accumulatornum_acc_stage = 1# Compute C stages (2 for double buffering)num_c_stage = 2- # Fixed overhead for barriers, etc.+ # Fixed overhead for barriersbarrier_bytes = 1024# Available SMEM for A/B/SF stagesavailable = smem_capacity // occupancy - c_bytes * num_c_stage - barrier_bytes- num_ab_stage = max(2, min(8, available // ab_sf_bytes_per_stage))+ num_ab_stage = max(2, min(30, 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.jitdef __call__(self,⋯ 6 unchanged linesmax_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- """+ """Execute the block-scaled GEMM operation."""m, n, k, l = shape_mnkl# Create tensors from pointers with proper layouts⋯ 5 unchanged linesb_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 is (M, N, L) row-majorc_layout = cute.make_layout((m, n, l), stride=(n, 1, m * n))c_tensor = cute.make_tensor(c_ptr, c_layout)⋯ 9 unchanged lines# Setup kernel attributesself._setup_attributes()- # Setup scale factor tensor layouts from pointers- # Scale factors use the tile_atom_to_shape_SF layout based on A/B shapes+ # Setup scale factor tensor layoutssfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, self.sf_vec_size)⋯ 102 unchanged linesself.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,+ # Compute grid using PersistentTileScheduler+ c_shape = cute.slice_(self.cta_tile_shape_mnk, (None, None, 0))+ gc = cute.zipped_divide(c_tensor, tiler=c_shape)+ num_ctas_mnl = gc[(0, (None, None, None))].shape+ cluster_shape_mnk = (*self.cluster_shape_mn, 1)++ tile_sched_params = utils.PersistentTileSchedulerParams(+ num_ctas_mnl, cluster_shape_mnk)+ grid = utils.StaticPersistentTileScheduler.get_grid_shape(+ tile_sched_params, max_active_clusters+ )- 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⋯ 51 unchanged linestma_tensor_sfb,tma_atom_c,tma_tensor_c,- c_tensor,self.cluster_layout_vmnk,self.cluster_layout_sfb_vmnk,self.a_smem_layout_staged,⋯ 2 unchanged linesself.sfb_smem_layout_staged,self.c_smem_layout_staged,self.epi_tile,- tiles_m,- tiles_n,- l,- total_work_items,+ tile_sched_params,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,- ),+ cluster=(*self.cluster_shape_mn, 1),min_blocks_per_mp=1,)⋯ 12 unchanged linesmSFB_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,⋯ 2 unchanged linessfb_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,+ tile_sched_params: utils.PersistentTileSchedulerParams,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.+ GPU device kernel for block-scaled GEMM with independent work loops per warp role."""- # 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.+ # TMA descriptor prefetchif warp_idx == self.tma_warp_id:cpasync.prefetch_descriptor(tma_atom_a)cpasync.prefetch_descriptor(tma_atom_b)⋯ 3 unchanged linesuse_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.+ # CTA/thread coordinatesbidx, 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+ mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)is_leader_cta = mma_tile_coord_v == 0cta_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)⋯ 1 unchanged linescta_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.+ # Shared memory allocationsmem = utils.SmemAllocator()storage = smem.allocate(self.shared_storage)- # Pipelines:- # - `PipelineTmaUmma` coordinates TMA producer(s) and UMMA consumer(s) via mbarrier arrays.+ # Initialize mainloop ab_pipelineab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1ab_pipeline_consumer_group = pipeline.CooperativeGroup(⋯ 9 unchanged linesdefer_sync=True,)- # Accumulator pipeline:- # - `PipelineUmmaAsync` synchronizes "accumulator stage" ownership across MMA and epilogue roles.+ # Initialize acc_pipelineacc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)num_acc_consumer_threads = len(self.epilog_warp_id) * (2 if use_2cta_instrs else 1⋯ 10 unchanged linesdefer_sync=True,)- # TMEM allocator:- # - Provides per-CTA allocations of Blackwell tensor memory columns for accumulators + SF tiles.+ # TMEM allocatortmem = utils.TmemAllocator(storage.tmem_holding_buf,barrier_for_retrieve=self.tmem_alloc_barrier,⋯ 18 unchanged linessSFA = 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=...)`.+ # TMA multicast masksa_full_mcast_mask = Noneb_full_mcast_mask = Nonesfa_full_mcast_mask = None⋯ 12 unchanged linescluster_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.+ # Global -> per-CTA tilesgA_mkl = cute.local_tile(mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None))⋯ 11 unchanged linesgC_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⋯ 5 unchanged linestCgSFB = 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.+ # TMA partitions for Aa_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape)⋯ 73 unchanged lines# Wait for cluster initpipeline_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)+ # =================================================================+ # INDEPENDENT WORK LOOPS PER WARP ROLE+ # No global synchronization barrier between work items+ # =================================================================- # 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- )+ # -----------------------------------------------------------------+ # TMA LOAD WARP (warp 5) - INDEPENDENT WORK LOOP+ # -----------------------------------------------------------------+ if warp_idx == self.tma_warp_id:+ # Create tile scheduler for TMA warp+ tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()+ )+ work_tile = tile_sched.initial_work_tile_info()- # 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,+ ab_producer_state = pipeline.make_pipeline_state(+ pipeline.PipelineUserType.Producer, self.num_ab_stage)- # -------------------------------------------------------------- # TMA LOAD WARP (warp 5)- # -------------------------------------------------------------- if warp_idx == self.tma_warp_id:+ while work_tile.is_valid_tile:+ # Get tile coord from tile scheduler+ cur_tile_coord = work_tile.tile_idx+ mma_tile_coord_mnl = (+ cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),+ cur_tile_coord[1],+ cur_tile_coord[2],+ )++ # Slice to per mma tile indextAgA_slice = tAgA[- (None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])+ (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]tBgB_slice = tBgB[- (None, mma_tile_coord_mnl[1], None, cluster_tile_coord_mnl[2])+ (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]tAgSFA_slice = tAgSFA[- (None, cluster_tile_coord_mnl[0], None, cluster_tile_coord_mnl[2])+ (None, mma_tile_coord_mnl[0], None, mma_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]+ 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])]+ slice_n = mma_tile_coord_mnl[1] // 2+ tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]+ # Reset and peek for first k tileab_producer_state.reset_count()peek_ab_empty_status = cutlass.Boolean(1)if ab_producer_state.count < k_tile_cnt:⋯ 1 unchanged linesab_producer_state)- for k_iter in cutlass.range(0, k_tile_cnt, 1, unroll=1): # noqa: B007+ # TMA load loop+ for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):+ # Wait for AB buffer emptyab_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.+ # TMA loadscute.copy(tma_atom_a,- tAgA_slice[(None, k_tile)],+ tAgA_slice[(None, ab_producer_state.count)],tAsA[(None, ab_producer_state.index)],tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=a_full_mcast_mask,)cute.copy(tma_atom_b,- tBgB_slice[(None, k_tile)],+ tBgB_slice[(None, ab_producer_state.count)],tBsB[(None, ab_producer_state.index)],tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=b_full_mcast_mask,)cute.copy(tma_atom_sfa,- tAgSFA_slice[(None, k_tile)],+ tAgSFA_slice[(None, ab_producer_state.count)],tAsSFA[(None, ab_producer_state.index)],tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=sfa_full_mcast_mask,)cute.copy(tma_atom_sfb,- tBgSFB_slice[(None, k_tile)],+ tBgSFB_slice[(None, ab_producer_state.count)],tBsSFB[(None, ab_producer_state.index)],tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=sfb_full_mcast_mask,)+ # Advance and peek for next k tileab_producer_state.advance()peek_ab_empty_status = cutlass.Boolean(1)if ab_producer_state.count < k_tile_cnt:⋯ 1 unchanged linesab_producer_state)- # -------------------------------------------------------------- # MMA WARP (warp 4)- # -------------------------------------------------------------- if warp_idx == self.mma_warp_id:+ # Advance to next tile+ tile_sched.advance_to_next_work()+ work_tile = tile_sched.get_current_work()++ # TMA warp tail+ ab_pipeline.producer_tail(ab_producer_state)++ # -----------------------------------------------------------------+ # MMA WARP (warp 4) - INDEPENDENT WORK LOOP+ # -----------------------------------------------------------------+ if warp_idx == self.mma_warp_id:+ # Wait for TMEM allocation+ tmem.wait_for_alloc()++ # Retrieve TMEM pointers+ acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)+ tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)++ # Make SFA tmem tensor+ 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)++ # Make SFB tmem tensor+ 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)++ # S2T copy partitions for scale factors+ (+ 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)++ # Create tile scheduler for MMA warp+ tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()+ )+ work_tile = tile_sched.initial_work_tile_info()++ ab_consumer_state = pipeline.make_pipeline_state(+ pipeline.PipelineUserType.Consumer, self.num_ab_stage+ )+ acc_producer_state = pipeline.make_pipeline_state(+ pipeline.PipelineUserType.Producer, self.num_acc_stage+ )++ while work_tile.is_valid_tile:+ # Get tile coord from tile scheduler+ cur_tile_coord = work_tile.tile_idx+ mma_tile_coord_mnl = (+ cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),+ cur_tile_coord[1],+ cur_tile_coord[2],+ )++ # Get accumulator stage indexif cutlass.const_expr(self.overlapping_accum):acc_stage_index = acc_producer_state.phase ^ 1else:⋯ 1 unchanged linestCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]+ # Peek for first k tileab_consumer_state.reset_count()peek_ab_full_status = cutlass.Boolean(1)if ab_consumer_state.count < k_tile_cnt and is_leader_cta:⋯ 1 unchanged linesab_consumer_state)+ # Wait for accumulator buffer emptyif is_leader_cta:acc_pipeline.producer_acquire(acc_producer_state)+ # Handle SFB offset for different tile shapestCtSFB_mma = tCtSFBif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):offset = (⋯ 20 unchanged lines)tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)+ # Reset accumulate flagtiled_mma.set(tcgen05.Field.ACCUMULATE, False)- for k_iter in range(k_tile_cnt): # noqa: B007+ # MMA mainloop+ for k_tile in range(k_tile_cnt):if is_leader_cta:+ # Wait for AB buffer fullab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status)+ # Copy SFA/SFB from smem to tmems2t_stage_coord = (None,None,⋯ 15 unchanged linestCtSFB_compact_s2t,)+ # MMA over k blocksnum_kblocks = cute.size(tCrA, mode=[2])for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):kblock_coord = (⋯ 11 unchanged linestcgen05.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,⋯ 4 unchanged linestiled_mma.set(tcgen05.Field.ACCUMULATE, True)+ # Release AB bufferab_pipeline.consumer_release(ab_consumer_state)+ # Peek for next k tileab_consumer_state.advance()peek_ab_full_status = cutlass.Boolean(1)if ab_consumer_state.count < k_tile_cnt:⋯ 2 unchanged linesab_consumer_state)+ # Commit accumulatorif 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:+ # Advance to next tile+ tile_sched.advance_to_next_work()+ work_tile = tile_sched.get_current_work()++ # MMA warp tail+ acc_pipeline.producer_tail(acc_producer_state)++ # -----------------------------------------------------------------+ # EPILOGUE WARPS (warps 0-3) - INDEPENDENT WORK LOOP+ # -----------------------------------------------------------------+ if warp_idx < self.mma_warp_id:+ # Allocate TMEM+ tmem.allocate(self.num_tmem_alloc_cols)++ # Wait for TMEM allocation+ tmem.wait_for_alloc()++ # Retrieve TMEM pointer+ acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)+ tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)++ # Partition for epilogue+ epi_tidx = tidx+ (+ tiled_copy_t2r,+ tTR_tAcc_base,+ tTR_rAcc,+ ) = self.epilog_tmem_copy_and_partition(+ epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs+ )++ tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)+ tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(+ tiled_copy_t2r, tTR_rC, epi_tidx, sC+ )+ (+ tma_atom_c,+ bSG_sC,+ bSG_gC_partitioned,+ ) = self.epilog_gmem_copy_and_partition(tma_atom_c, tCgC, epi_tile, sC)++ # Create tile scheduler for Epilogue warps+ tile_sched = utils.StaticPersistentTileScheduler.create(+ tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()+ )+ work_tile = tile_sched.initial_work_tile_info()++ acc_consumer_state = pipeline.make_pipeline_state(+ pipeline.PipelineUserType.Consumer, self.num_acc_stage+ )++ # C TMA store pipeline+ c_producer_group = pipeline.CooperativeGroup(+ pipeline.Agent.Thread,+ 32 * len(self.epilog_warp_id),+ )+ c_pipeline = pipeline.PipelineTmaStore.create(+ num_stages=self.num_c_stage,+ producer_group=c_producer_group,+ )++ while work_tile.is_valid_tile:+ # Get tile coord from tile scheduler+ cur_tile_coord = work_tile.tile_idx+ mma_tile_coord_mnl = (+ cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),+ cur_tile_coord[1],+ cur_tile_coord[2],+ )++ # Slice output tensorbSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]+ # Get accumulator stage indexif cutlass.const_expr(self.overlapping_accum):acc_stage_index = acc_consumer_state.phase+ reverse_subtile = (+ cutlass.Boolean(True)+ if acc_stage_index == 0+ else cutlass.Boolean(False)+ )else:acc_stage_index = acc_consumer_state.index⋯ 1 unchanged lines(None, None, None, None, None, acc_stage_index)]+ # Wait for accumulator buffer fullacc_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))+ # Store accumulator to global memory in subtilessubtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])- num_prev_subtiles = 0+ num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt- if warp_idx == self.epilog_warp_id[0] and subtile_cnt > 0:- c_pipeline.producer_acquire()+ for subtile_idx in cutlass.range(subtile_cnt):+ real_subtile_idx = subtile_idx+ if cutlass.const_expr(self.overlapping_accum):+ if reverse_subtile:+ real_subtile_idx = (+ self.cta_tile_shape_mnk[1] // self.epi_tile_n+ - 1+ - subtile_idx+ )- 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)+ # Load accumulator from TMEM to registers+ tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]+ cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)- 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+ # Early release for overlapping accum+ if cutlass.const_expr(self.overlapping_accum):+ if subtile_idx == self.iter_acc_early_release_in_epilogue:+ cute.arch.fence_view_async_tmem_load()+ with cute.arch.elect_one():+ acc_pipeline.consumer_release(acc_consumer_state)+ acc_consumer_state.advance()- 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)+ # Convert to C type+ acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()+ acc_vec = epilogue_op(acc_vec.to(self.c_dtype))+ tRS_rC.store(acc_vec)- 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+ # Store C to shared memory+ c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stagecute.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).+ # Fence and barriercute.arch.fence_proxy(cute.arch.ProxyKind.async_shared,space=cute.arch.SharedSpace.shared_cta,)self.epilog_sync_barrier.arrive_and_wait()+ # TMA store Cif warp_idx == self.epilog_warp_id[0]:cute.copy(tma_atom_c,bSG_sC[(None, c_buffer)],- bSG_gC[(None, subtile_idx)],+ bSG_gC[(None, real_subtile_idx)],⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON