submission 179783
rcmalli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1793 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-179783?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:9a3a9081535360164ebece48eb3dfcc2dd9def9389166f9852b9324393234343
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.py1793 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
# NOTE: TMEM columns must always be computed with factor 4, NOT mma_inst_tile_k.
# The SMEM layout shape is constant (due to divide-by-mma_inst_tile_k cancellation),
# so TMEM layout derived from SMEM always expects the same number of columns.
# Using mma_inst_tile_k here causes buffer mismatch when mma_inst_tile_k != 4.
sf_atom_mn = 32
tmem_k_factor = 4 # Hardware constant, not mma_inst_tile_k
# Double-buffer scale factors in TMEM for S2T pipelining
self.num_sf_stages = 2
self.num_sfa_tmem_cols_per_stage = (
self.cta_tile_shape_mnk[0] // sf_atom_mn
) * tmem_k_factor
self.num_sfb_tmem_cols_per_stage = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
) * tmem_k_factor
# Total SF columns = (SFA + SFB) * num_stages
self.num_sf_tmem_cols = (
self.num_sfa_tmem_cols_per_stage + self.num_sfb_tmem_cols_per_stage
) * self.num_sf_stages
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 overlapping accumulator only for wide N tiles (matches reference kernels).
num_acc_stage = 1 if mma_tiler[1] == 256 else 2
# 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)
# Create TMEM layouts for scale factors
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)),
)
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)),
)
# Double-buffered SFA TMEM tensors (buffer 0 and buffer 1)
# Layout: [Acc] [SFA_0] [SFA_1] [SFB_0] [SFB_1]
# Default to manual offsets to keep overlap layout invariants intact.
sfa_tmem_ptr_0 = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols,
dtype=self.sf_dtype,
)
sfa_tmem_ptr_1 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols_per_stage,
dtype=self.sf_dtype,
)
tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)
sfb_base_offset = (
self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols_per_stage * self.num_sf_stages
)
sfb_tmem_ptr_0 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset,
dtype=self.sf_dtype,
)
sfb_tmem_ptr_1 = cute.recast_ptr(
acc_tmem_ptr
+ sfb_base_offset
+ self.num_sfb_tmem_cols_per_stage,
dtype=self.sf_dtype,
)
tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)
sfb_tmem_cols = self.num_sfb_tmem_cols_per_stage
# Non-overlapping path can safely use layout-derived offsets.
if cutlass.const_expr(not self.overlapping_accum):
acc_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
sfa_tmem_ptr_0 = cute.recast_ptr(
acc_tmem_ptr + acc_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)
sfa_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFA_0)
sfa_tmem_ptr_1 = cute.recast_ptr(
acc_tmem_ptr + acc_tmem_cols + sfa_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)
sfb_base_offset = acc_tmem_cols + sfa_tmem_cols * self.num_sf_stages
sfb_tmem_ptr_0 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset,
dtype=self.sf_dtype,
)
tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)
sfb_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFB_0)
sfb_tmem_ptr_1 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)
# S2T copy partitions for scale factors (both buffers)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t_0,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_0)
(
_,
_,
tCtSFA_compact_s2t_1,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_1)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t_0,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_0)
(
_,
_,
tCtSFB_compact_s2t_1,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_1)
# 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 (double-buffered)
# Create SFB MMA tensors for both buffers with appropriate offsets
tCtSFB_mma_0 = tCtSFB_0
tCtSFB_mma_1 = tCtSFB_1
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_0 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset + offset,
dtype=self.sf_dtype,
)
shifted_ptr_1 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
dtype=self.sf_dtype,
)
tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, 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_0 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset + offset,
dtype=self.sf_dtype,
)
shifted_ptr_1 = cute.recast_ptr(
acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,
dtype=self.sf_dtype,
)
tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)
tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, tCtSFB_layout)
# Reset accumulate flag
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
# MMA mainloop with double-buffered SF TMEM
for k_tile in range(k_tile_cnt):
# Default to buffer 0 so variables exist outside dynamic control flow.
tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
tCtSFA_mma = tCtSFA_0
tCtSFB_mma = tCtSFB_mma_0
if is_leader_cta:
# Select SF TMEM buffer based on pipeline stage parity
sf_stage_is_zero = (ab_consumer_state.index & 1) == 0
if sf_stage_is_zero:
tCtSFA_compact_s2t = tCtSFA_compact_s2t_0
tCtSFB_compact_s2t = tCtSFB_compact_s2t_0
tCtSFA_mma = tCtSFA_0
tCtSFB_mma = tCtSFB_mma_0
else:
tCtSFA_compact_s2t = tCtSFA_compact_s2t_1
tCtSFB_compact_s2t = tCtSFB_compact_s2t_1
tCtSFA_mma = tCtSFA_1
tCtSFB_mma = tCtSFB_mma_1
# Wait for AB buffer full
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
# Copy SFA/SFB from smem to tmem (using buffer 0)
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_mma[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=3, mma_inst_tile_k=4
),
(128, 4096, 7168, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
),
(128, 7168, 2048, 1): KernelConfig(
(128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4
),
}
# Default fallback configuration for unseen shapes.
# mma_inst_tile_k=2 gives smaller K-tiles (128 vs 256), allowing more AB stages
_DEFAULT_CONFIG = KernelConfig(
(128, 64),
(1, 1),
num_ab_stage=10,
num_c_stage=2,
mma_inst_tile_k=2,
)
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 · 1793 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 168494.
⋯ 352 unchanged linesself.overlapping_accum = self.num_acc_stage == 1# TMEM column counts+ # NOTE: TMEM columns must always be computed with factor 4, NOT mma_inst_tile_k.+ # The SMEM layout shape is constant (due to divide-by-mma_inst_tile_k cancellation),+ # so TMEM layout derived from SMEM always expects the same number of columns.+ # Using mma_inst_tile_k here causes buffer mismatch when mma_inst_tile_k != 4.sf_atom_mn = 32- self.num_sfa_tmem_cols = (+ tmem_k_factor = 4 # Hardware constant, not mma_inst_tile_k+ # Double-buffer scale factors in TMEM for S2T pipelining+ self.num_sf_stages = 2+ self.num_sfa_tmem_cols_per_stage = (self.cta_tile_shape_mnk[0] // sf_atom_mn- ) * self.mma_inst_tile_k- self.num_sfb_tmem_cols = (+ ) * tmem_k_factor+ self.num_sfb_tmem_cols_per_stage = (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+ ) * tmem_k_factor+ # Total SF columns = (SFA + SFB) * num_stages+ self.num_sf_tmem_cols = (+ self.num_sfa_tmem_cols_per_stage + self.num_sfb_tmem_cols_per_stage+ ) * self.num_sf_stagesself.num_accumulator_tmem_cols = (self.cta_tile_shape_mnk[1] * self.num_acc_stageif not self.overlapping_accum⋯ 38 unchanged linesab_sf_bytes_per_stage = a_bytes + b_bytes + sfa_bytes + sfb_bytes- # Use num_acc_stage=1 for overlapping accumulator- num_acc_stage = 1+ # Use overlapping accumulator only for wide N tiles (matches reference kernels).+ num_acc_stage = 1 if mma_tiler[1] == 256 else 2# Compute C stages (2 for double buffering)num_c_stage = 2⋯ 605 unchanged linesacc_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,- )+ # Create TMEM layouts for scale factorstCtSFA_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+ # Double-buffered SFA TMEM tensors (buffer 0 and buffer 1)+ # Layout: [Acc] [SFA_0] [SFA_1] [SFB_0] [SFB_1]+ # Default to manual offsets to keep overlap layout invariants intact.+ sfa_tmem_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + self.num_accumulator_tmem_cols,+ dtype=self.sf_dtype,+ )+ sfa_tmem_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr+ + self.num_accumulator_tmem_cols+ + self.num_sfa_tmem_cols_per_stage,+ dtype=self.sf_dtype,+ )+ tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)+ tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)++ sfb_base_offset = (+ self.num_accumulator_tmem_cols+ + self.num_sfa_tmem_cols_per_stage * self.num_sf_stages+ )+ sfb_tmem_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset,+ dtype=self.sf_dtype,+ )+ sfb_tmem_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr+ + sfb_base_offset+ + self.num_sfb_tmem_cols_per_stage,+ dtype=self.sf_dtype,+ )+ tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)+ tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)+ sfb_tmem_cols = self.num_sfb_tmem_cols_per_stage++ # Non-overlapping path can safely use layout-derived offsets.+ if cutlass.const_expr(not self.overlapping_accum):+ acc_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)+ sfa_tmem_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + acc_tmem_cols,+ dtype=self.sf_dtype,+ )+ tCtSFA_0 = cute.make_tensor(sfa_tmem_ptr_0, tCtSFA_layout)+ sfa_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFA_0)+ sfa_tmem_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr + acc_tmem_cols + sfa_tmem_cols,+ dtype=self.sf_dtype,+ )+ tCtSFA_1 = cute.make_tensor(sfa_tmem_ptr_1, tCtSFA_layout)++ sfb_base_offset = acc_tmem_cols + sfa_tmem_cols * self.num_sf_stages+ sfb_tmem_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset,+ dtype=self.sf_dtype,+ )+ tCtSFB_0 = cute.make_tensor(sfb_tmem_ptr_0, tCtSFB_layout)+ sfb_tmem_cols = tcgen05.find_tmem_tensor_col_offset(tCtSFB_0)+ sfb_tmem_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols,+ dtype=self.sf_dtype,+ )+ tCtSFB_1 = cute.make_tensor(sfb_tmem_ptr_1, tCtSFB_layout)++ # S2T copy partitions for scale factors (both buffers)(tiled_copy_s2t_sfa,tCsSFA_compact_s2t,- tCtSFA_compact_s2t,- ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)+ tCtSFA_compact_s2t_0,+ ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_0)(+ _,+ _,+ tCtSFA_compact_s2t_1,+ ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA_1)+ (tiled_copy_s2t_sfb,tCsSFB_compact_s2t,- tCtSFB_compact_s2t,- ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)+ tCtSFB_compact_s2t_0,+ ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_0)+ (+ _,+ _,+ tCtSFB_compact_s2t_1,+ ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB_1)# Create tile scheduler for MMA warptile_sched = utils.StaticPersistentTileScheduler.create(⋯ 37 unchanged linesif is_leader_cta:acc_pipeline.producer_acquire(acc_producer_state)- # Handle SFB offset for different tile shapes- tCtSFB_mma = tCtSFB+ # Handle SFB offset for different tile shapes (double-buffered)+ # Create SFB MMA tensors for both buffers with appropriate offsets+ tCtSFB_mma_0 = tCtSFB_0+ tCtSFB_mma_1 = tCtSFB_1if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):offset = (cutlass.Int32(2)if mma_tile_coord_mnl[1] % 2 == 1else cutlass.Int32(0))- shifted_ptr = cute.recast_ptr(- acc_tmem_ptr- + self.num_accumulator_tmem_cols- + self.num_sfa_tmem_cols- + offset,+ shifted_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset + offset,dtype=self.sf_dtype,)- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)+ shifted_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,+ dtype=self.sf_dtype,+ )+ tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)+ tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, 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,+ shifted_ptr_0 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset + offset,dtype=self.sf_dtype,)- tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)+ shifted_ptr_1 = cute.recast_ptr(+ acc_tmem_ptr + sfb_base_offset + sfb_tmem_cols + offset,+ dtype=self.sf_dtype,+ )+ tCtSFB_mma_0 = cute.make_tensor(shifted_ptr_0, tCtSFB_layout)+ tCtSFB_mma_1 = cute.make_tensor(shifted_ptr_1, tCtSFB_layout)# Reset accumulate flagtiled_mma.set(tcgen05.Field.ACCUMULATE, False)- # MMA mainloop+ # MMA mainloop with double-buffered SF TMEMfor k_tile in range(k_tile_cnt):+ # Default to buffer 0 so variables exist outside dynamic control flow.+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_0+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_0+ tCtSFA_mma = tCtSFA_0+ tCtSFB_mma = tCtSFB_mma_0+if is_leader_cta:+ # Select SF TMEM buffer based on pipeline stage parity+ sf_stage_is_zero = (ab_consumer_state.index & 1) == 0+ if sf_stage_is_zero:+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_0+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_0+ tCtSFA_mma = tCtSFA_0+ tCtSFB_mma = tCtSFB_mma_0+ else:+ tCtSFA_compact_s2t = tCtSFA_compact_s2t_1+ tCtSFB_compact_s2t = tCtSFB_compact_s2t_1+ tCtSFA_mma = tCtSFA_1+ tCtSFB_mma = tCtSFB_mma_1+# Wait for AB buffer fullab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status)- # Copy SFA/SFB from smem to tmem+ # Copy SFA/SFB from smem to tmem (using buffer 0)s2t_stage_coord = (None,None,⋯ 27 unchanged linessf_kblock_coord = (None, None, kblock_idx)tiled_mma.set(- tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator+ tcgen05.Field.SFA, tCtSFA_mma[sf_kblock_coord].iterator)tiled_mma.set(tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator⋯ 299 unchanged lines_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, 64), (1, 2), num_ab_stage=7, num_c_stage=3, 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, 64), (1, 2), num_ab_stage=7, num_c_stage=3, 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+ (128, 64), (1, 2), num_ab_stage=7, num_c_stage=3, mma_inst_tile_k=4),}# Default fallback configuration for unseen shapes.+ # mma_inst_tile_k=2 gives smaller K-tiles (128 vs 256), allowing more AB stages_DEFAULT_CONFIG = KernelConfig((128, 64),(1, 1),- num_ab_stage=7,+ num_ab_stage=10,num_c_stage=2,- mma_inst_tile_k=4,+ mma_inst_tile_k=2,)
scrolls · 289 diff lines total
Best evidence level for this revision: reported
JSON