submission 179024
currybab · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1425 lines, June 9 Researcher Reciprocity License v1.0.
submission_8.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-179024?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:ba9c07a556c4e9685381a527ec52acb100ffab6aa725177be63e277eb992472c
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
self.cta_sync_barrier = pipeline.NamedBarrier(persistent-kernel
"""Persistent block-scaled GEMM kernel for SM100 (Blackwell).shared-memory
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")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_8.py1425 lines
from typing import Type, Tuple, Union, Dict
import inspect
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.pipeline as pipeline
import cutlass.utils as utils
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
# ---------------------------------------------------------------------------
# Optional L2 cache eviction priority hint (safe no-op if unsupported).
#
# Motivation for the 3 target cases (m=128, l=1):
# - A and SFA are reused across many N-tiles (each CTA reads the same A/SFA).
# - B/SFB and C are mostly
# So we *prefer* to keep A/SFA in L2 longer and let B/SFB/C evict first.
# ---------------------------------------------------------------------------
_HAS_COPY_CACHE_POLICY = False
_EVICT_FIRST = None
_EVICT_LAST = None
try:
_sig = inspect.signature(cute.copy)
_HAS_COPY_CACHE_POLICY = "cache_policy" in _sig.parameters
except Exception:
_HAS_COPY_CACHE_POLICY = False
try:
_EVICT_FIRST = cute.CacheEvictionPriority.EVICT_FIRST
_EVICT_LAST = cute.CacheEvictionPriority.EVICT_LAST
except Exception:
_EVICT_FIRST = None
_EVICT_LAST = None
_CACHE_A = _EVICT_LAST
_CACHE_SFA = _EVICT_LAST
_CACHE_B = _EVICT_FIRST
_CACHE_SFB = _EVICT_FIRST
_CACHE_C = _EVICT_FIRST
def _as_int64_cache_policy(policy):
"""Convert cache-policy enum/int into the Int64 object expected by cpasync/TMA.
Some CUTLASS-DSL versions require `cache_policy` to be a `cutlass.Int64`.
Passing a Python enum (like `cute.CacheEvictionPriority`) triggers:
ValueError: expects `Int64` value to be provided via the cache_policy kw argument
Return None if conversion fails, in which case we simply omit the argument.
"""
if policy is None:
return None
# Already the correct DSL scalar?
try:
if isinstance(policy, cutlass.Int64):
return policy
except Exception:
pass
# IntEnum / int-like
try:
return cutlass.Int64(int(policy))
except Exception:
pass
# Enum with `.value`
try:
return cutlass.Int64(int(getattr(policy, "value")))
except Exception:
return None
def _cute_copy_with_cache_policy(atom, src, dst, cache_policy, **kwargs):
"""cute.copy wrapper that *optionally* adds cache_policy=... when supported."""
if _HAS_COPY_CACHE_POLICY and (cache_policy is not None):
cp_i64 = _as_int64_cache_policy(cache_policy)
if cp_i64 is not None:
try:
return cute.copy(atom, src, dst, cache_policy=cp_i64, **kwargs)
except (TypeError, ValueError):
# Older DSL / op variant without cache_policy kw, or wrong scalar type
pass
return cute.copy(atom, src, dst, **kwargs)
class Sm100BlockScaledPersistentDenseGemmKernel:
"""Persistent block-scaled GEMM kernel for SM100 (Blackwell).
C = A x SFA x B x SFB (block-scaled NVF4 path)
Notes:
- A/B: Float4E2M1FN
- SF : Float8E4M3FN (vector size 16)
- Acc: Float32
- C : Float16/BFloat16/Float32
- A/B are K-major in this challenge setting.
"""
def __init__(
self,
sf_vec_size: int,
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
mma_inst_tile_k: int = 4,
):
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
# K dimension is deferred in _setup_attributes
self.mma_tiler = (*mma_tiler_mn, 1)
# Override K-tiling (# of MMA-Inst-K blocks per CTA tile)
self.mma_inst_tile_k = mma_inst_tile_k
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.occupancy = 1
# Warp specialization
self.epilog_warp_id = (0, 1, 2, 3)
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_cta = 32 * len(
(self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
)
# Named barriers
self.cta_sync_barrier = pipeline.NamedBarrier(
barrier_id=1, num_threads=self.threads_per_cta
)
self.epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=2, num_threads=32 * len(self.epilog_warp_id)
)
self.tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=3,
num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
)
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
# TMEM columns capacity (Blackwell)
self.num_tmem_alloc_cols = 512
def _setup_attributes(self):
# MMA instruction shapes
self.mma_inst_shape_mn = (self.mma_tiler[0], self.mma_tiler[1])
self.mma_inst_shape_mn_sfb = (
self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
cute.round_up(self.mma_inst_shape_mn[1], 128),
)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
mma_inst_tile_k = self.mma_inst_tile_k
self.mma_tiler = (
self.mma_inst_shape_mn[0],
self.mma_inst_shape_mn[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.mma_tiler_sfb = (
self.mma_inst_shape_mn_sfb[0],
self.mma_inst_shape_mn_sfb[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.cta_tile_shape_mnk = (
self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler[1],
self.mma_tiler[2],
)
# 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 CTAs
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 subtile
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
self.cta_tile_shape_mnk,
self.use_2cta_instrs,
self.c_layout,
self.c_dtype,
)
# 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,
)
# SMEM layouts
self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(
tiled_mma,
self.mma_tiler,
self.a_dtype,
self.num_ab_stage,
)
self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(
tiled_mma,
self.mma_tiler,
self.b_dtype,
self.num_ab_stage,
)
self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
)
self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
)
self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
self.c_dtype,
self.c_layout,
self.epi_tile,
self.num_c_stage,
)
@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,
layouts: cutlass.Constexpr[
Tuple[tcgen05.OperandMajorMode, tcgen05.OperandMajorMode, utils.LayoutEnum]
],
problem_mnkl: Tuple[int, int, int, int],
max_active_clusters: cutlass.Constexpr,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
# Static attrs
self.a_dtype: Type[cutlass.Numeric] = a_ptr.value_type
self.b_dtype: Type[cutlass.Numeric] = b_ptr.value_type
self.sf_dtype: Type[cutlass.Numeric] = sfa_ptr.value_type
self.c_dtype: Type[cutlass.Numeric] = c_ptr.value_type
m, n, k, l = problem_mnkl
self.a_major_mode, self.b_major_mode, self.c_layout = layouts
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")
self._setup_attributes()
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr,
cute.make_layout((cute.assume(m, 32), n, l), stride=(n, 1, m * n)),
)
# SF tensors
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)
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,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
# TMA load A
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,
)
# TMA load B
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,
)
# TMA load SFA
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,
)
# TMA load SFB
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,
)
# SFB slicing hack for N=192 (keep)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x = tma_tensor_sfb.stride[0][1]
y = cute.ceil_div(tma_tensor_sfb.shape[0][1], 4)
new_shape = (
(tma_tensor_sfb.shape[0][0], ((2, 2), y)),
tma_tensor_sfb.shape[1],
tma_tensor_sfb.shape[2],
)
x_times_3 = 3 * x
new_stride = (
(tma_tensor_sfb.stride[0][0], ((x, x), x_times_3)),
tma_tensor_sfb.stride[1],
tma_tensor_sfb.stride[2],
)
tma_tensor_sfb_new_layout = cute.make_layout(new_shape, stride=new_stride)
tma_tensor_sfb = cute.make_tensor(
tma_tensor_sfb.iterator, tma_tensor_sfb_new_layout
)
# Tx bytes
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 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,
)
# Grid
self.tile_sched_params, grid = self._compute_grid(
c_tensor,
self.cta_tile_shape_mnk,
self.cluster_shape_mn,
max_active_clusters,
)
self.buffer_align_bytes = 1024
@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,
self.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,
)
return
@cute.kernel
def kernel(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a: cute.CopyAtom,
mA_mkl: cute.Tensor,
tma_atom_b: cute.CopyAtom,
mB_nkl: cute.Tensor,
tma_atom_sfa: cute.CopyAtom,
mSFA_mkl: cute.Tensor,
tma_atom_sfb: cute.CopyAtom,
mSFB_nkl: cute.Tensor,
tma_atom_c: cute.CopyAtom,
mC_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
tile_sched_params: utils.PersistentTileSchedulerParams,
epilogue_op: cutlass.Constexpr,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# Prefetch TMA desc
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
# Block/cluster coords
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()[0]
# Allocate smem storage
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
# 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,
)
# 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,
)
# 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
if cute.size(self.cluster_shape_mn) > 1:
cute.arch.cluster_arrive_relaxed()
# SMEM tensors
sC = storage.sC.get_tensor(
c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
)
sA = storage.sA.get_tensor(
a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
)
sB = storage.sB.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
# Multicast masks
a_full_mcast_mask = 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,
)
# Local 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])
# MMA partitions
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgB = thr_mma.partition_B(gB_nkl)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
# TMA partitions
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, tAgA = cpasync.tma_partition(
tma_atom_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, tBgB = cpasync.tma_partition(
tma_atom_b,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB, 0, 3),
)
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cute.nvgpu.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)
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, tBgSFB = cute.nvgpu.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)
# Fragments
tCrA = tiled_mma.make_fragment_A(sA)
tCrB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
# Cluster wait before TMEM alloc
if cute.size(self.cluster_shape_mn) > 1:
cute.arch.cluster_wait()
else:
self.cta_sync_barrier.arrive_and_wait()
# ------------------ TMA warp ------------------
if warp_idx == self.tma_warp_id:
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:
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],
)
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])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
for _k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
_cute_copy_with_cache_policy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
cache_policy=_CACHE_A,
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
)
_cute_copy_with_cache_policy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
cache_policy=_CACHE_B,
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
)
_cute_copy_with_cache_policy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
cache_policy=_CACHE_SFA,
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
)
_cute_copy_with_cache_policy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
cache_policy=_CACHE_SFB,
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
)
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
ab_pipeline.producer_tail(ab_producer_state)
# ------------------ MMA warp ------------------
if warp_idx == self.mma_warp_id:
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 + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base),
dtype=self.sf_dtype,
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
sfb_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA),
dtype=self.sf_dtype,
)
tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
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:
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],
)
tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
ab_consumer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
# TMEM pointer offset hacks (keep)
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
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
+ 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
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
# Reset ACCUMULATE each output tile
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
# Hoist num_kblocks
num_kblocks = cute.size(tCrA, mode=[2])
for k_tile in range(k_tile_cnt):
if is_leader_cta:
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
s2t_stage_coord = (
None,
None,
None,
None,
ab_consumer_state.index,
)
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,
)
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,
)
# Set ACCUMULATE=True exactly once per output tile
if k_tile == 0:
if kblock_idx == 0:
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
acc_pipeline.producer_tail(acc_producer_state)
# ------------------ Epilogue warps ------------------
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_rAcc,
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
(
tma_atom_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
)
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_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:
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],
)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_consumer_state.index)
]
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))
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):
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
_cute_copy_with_cache_policy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, subtile_idx)],
cache_policy=_CACHE_C,
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
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: cute.Tensor,
tSF: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
tCsSF_compact = cute.filter_zeros(sSF)
tCtSF_compact = cute.filter_zeros(tSF)
copy_atom_s2t = cute.make_copy_atom(
tcgen05.Cp4x32x128bOp(self.cta_group),
self.sf_dtype,
)
tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
thr_copy_s2t = tiled_copy_s2t.get_slice(0)
tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t, tCsSF_compact_s2t_
)
tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t
def epilog_tmem_copy_and_partition(
self,
tidx: cutlass.Int32,
tAcc: cute.Tensor,
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
use_2cta_instrs: Union[cutlass.Boolean, bool],
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
tAcc_epi = cute.flat_divide(
tAcc[((None, None), 0, 0, None)],
epi_tile,
)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
gC_mnl_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
tTR_rAcc = cute.make_rmem_tensor(
tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
return tiled_copy_t2r, tTR_tAcc, tTR_rAcc
def epilog_smem_copy_and_partition(
self,
tiled_copy_t2r: cute.TiledCopy,
tTR_rC: cute.Tensor,
tidx: cutlass.Int32,
sC: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
copy_atom_r2s = sm100_utils.get_smem_store_op(
self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
)
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(
self,
tidx: cutlass.Int32,
atom: Union[cute.CopyAtom, cute.TiledCopy],
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
gC_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tma_atom_c = atom
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
@staticmethod
def _compute_stages(
tiled_mma: cute.TiledMma,
mma_tiler_mnk: Tuple[int, int, int],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
epi_tile: cute.Tile,
c_dtype: Type[cutlass.Numeric],
c_layout: utils.LayoutEnum,
sf_dtype: Type[cutlass.Numeric],
sf_vec_size: int,
smem_capacity: int,
occupancy: int,
) -> Tuple[int, int, int]:
# ACC stages (restore original heuristic: N==256 -> 1, else 2)
num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2
num_c_stage = 2
a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
tiled_mma, mma_tiler_mnk, a_dtype, 1
)
b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
tiled_mma, mma_tiler_mnk, b_dtype, 1
)
sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
tiled_mma, mma_tiler_mnk, sf_vec_size, 1
)
sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
tiled_mma, mma_tiler_mnk, sf_vec_size, 1
)
c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(
c_dtype, c_layout, epi_tile, 1
)
ab_bytes_per_stage = (
cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
+ cute.size_in_bytes(b_dtype, b_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
)
mbar_helpers_bytes = 1024
c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one)
c_bytes = c_bytes_per_stage * num_c_stage
num_ab_stage = (
smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
) // ab_bytes_per_stage
num_c_stage += (
smem_capacity
- occupancy * ab_bytes_per_stage * num_ab_stage
- occupancy * (mbar_helpers_bytes + c_bytes)
) // (occupancy * c_bytes_per_stage)
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
def _compute_grid(
c: cute.Tensor,
cta_tile_shape_mnk: Tuple[int, int, int],
cluster_shape_mn: Tuple[int, int],
max_active_clusters: cutlass.Constexpr,
) -> Tuple[utils.PersistentTileSchedulerParams, Tuple[int, int, int]]:
c_shape = cute.slice_(cta_tile_shape_mnk, (None, None, 0))
gc = cute.zipped_divide(c, tiler=c_shape)
num_ctas_mnl = gc[(0, (None, None, None))].shape
cluster_shape_mnl = (*cluster_shape_mn, 1)
tile_sched_params = utils.PersistentTileSchedulerParams(
num_ctas_mnl, cluster_shape_mnl
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
return tile_sched_params, grid
# ----------------- Hackathon wrapper -----------------
# Data types
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
# Default configuration (baseline)
_SF_VEC_SIZE = 16
_MMA_TILER_MN = (128, 64)
_CLUSTER_SHAPE_MN = (1, 1)
# Compile options
_COMPILE_OPT_LEVEL = 3
# Single supported K-tiling (avoid compile-fail + fallback noise)
_GEMM = Sm100BlockScaledPersistentDenseGemmKernel(
sf_vec_size=_SF_VEC_SIZE,
mma_tiler_mn=_MMA_TILER_MN,
cluster_shape_mn=_CLUSTER_SHAPE_MN,
mma_inst_tile_k=4,
)
# Compile cache keyed by problem_size (m,n,k,l)
_compiled: Dict[Tuple[int, int, int, int], object] = {}
# NOTE: Do NOT query HardwareInfo at import time.
#
# The benchmark harness imports this module inside multiprocessing workers.
# At import time there may be no active CUDA context, and `utils.HardwareInfo()`
# uses CUDA driver APIs that require a valid current context. That can trigger:
# CUDA_ERROR_INVALID_CONTEXT
#
# For our fixed cluster shape (1,1) and the three target problems, the total
# number of output tiles is <= 112, so as long as max_active_clusters is >= 112
# it will not cap the grid. We therefore use a conservative constant to avoid
# any driver calls during import.
_MAX_ACTIVE_CLUSTERS = 1024
def compile_kernel(problem_size: Tuple[int, int, int, int]):
"""Compile and cache a kernel specialized for (m,n,k,l)."""
ps = tuple(problem_size)
if ps in _compiled:
return _compiled[ps]
# Dummy pointers for compilation
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
# A/B are K-major for this task; output is row-major.
layouts = (
tcgen05.OperandMajorMode.K,
tcgen05.OperandMajorMode.K,
utils.LayoutEnum.ROW_MAJOR,
)
compiled = cute.compile(
_GEMM,
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
layouts,
ps,
_MAX_ACTIVE_CLUSTERS,
options=f"--opt-level {_COMPILE_OPT_LEVEL}",
)
_compiled[ps] = compiled
return compiled
def custom_kernel(data: input_t) -> output_t:
a, b, _, _, sfa_permuted, sfb_permuted, c = data
m, k_packed, l = a.shape
n, _, _ = b.shape
# Torch packs 2 FP4 values into one byte (e2m1_x2), so logical K doubles
k = k_packed * 2
problem_size = (m, n, k, l)
compiled = compile_kernel(problem_size)
a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(
sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
)
sfb_ptr = make_ptr(
sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
)
c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
# Constexpr args (layouts/max_active_clusters) are baked in by cute.compile.
compiled(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, problem_size)
return c
scrolls · 1425 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 169378.
- # submission_cluster11_newchallenge_ktile8_dispatch.py- from typing import Type, Tuple, Union+ from typing import Type, Tuple, Union, Dict+ import inspect- import torch-import cutlassimport cutlass.cute as cutefrom cutlass.cute.nvgpu import cpasync, tcgen05- import cutlass.utils as utilsimport cutlass.pipeline as pipeline+ import cutlass.utils as utilsimport cutlass.utils.blackwell_helpers as sm100_utilsimport cutlass.utils.blockscaled_layout as blockscaled_utilsfrom cutlass.cute.runtime import make_ptr+from task import input_t, output_t+ # ---------------------------------------------------------------------------+ # Optional L2 cache eviction priority hint (safe no-op if unsupported).+ #+ # Motivation for the 3 target cases (m=128, l=1):+ # - A and SFA are reused across many N-tiles (each CTA reads the same A/SFA).+ # - B/SFB and C are mostly+ # So we *prefer* to keep A/SFA in L2 longer and let B/SFB/C evict first.+ # ---------------------------------------------------------------------------+ _HAS_COPY_CACHE_POLICY = False+ _EVICT_FIRST = None+ _EVICT_LAST = None++ try:+ _sig = inspect.signature(cute.copy)+ _HAS_COPY_CACHE_POLICY = "cache_policy" in _sig.parameters+ except Exception:+ _HAS_COPY_CACHE_POLICY = False++ try:+ _EVICT_FIRST = cute.CacheEvictionPriority.EVICT_FIRST+ _EVICT_LAST = cute.CacheEvictionPriority.EVICT_LAST+ except Exception:+ _EVICT_FIRST = None+ _EVICT_LAST = None++ _CACHE_A = _EVICT_LAST+ _CACHE_SFA = _EVICT_LAST+ _CACHE_B = _EVICT_FIRST+ _CACHE_SFB = _EVICT_FIRST+ _CACHE_C = _EVICT_FIRST+++ def _as_int64_cache_policy(policy):+ """Convert cache-policy enum/int into the Int64 object expected by cpasync/TMA.++ Some CUTLASS-DSL versions require `cache_policy` to be a `cutlass.Int64`.+ Passing a Python enum (like `cute.CacheEvictionPriority`) triggers:+ ValueError: expects `Int64` value to be provided via the cache_policy kw argument++ Return None if conversion fails, in which case we simply omit the argument.+ """+ if policy is None:+ return None+ # Already the correct DSL scalar?+ try:+ if isinstance(policy, cutlass.Int64):+ return policy+ except Exception:+ pass++ # IntEnum / int-like+ try:+ return cutlass.Int64(int(policy))+ except Exception:+ pass++ # Enum with `.value`+ try:+ return cutlass.Int64(int(getattr(policy, "value")))+ except Exception:+ return None+++ def _cute_copy_with_cache_policy(atom, src, dst, cache_policy, **kwargs):+ """cute.copy wrapper that *optionally* adds cache_policy=... when supported."""+ if _HAS_COPY_CACHE_POLICY and (cache_policy is not None):+ cp_i64 = _as_int64_cache_policy(cache_policy)+ if cp_i64 is not None:+ try:+ return cute.copy(atom, src, dst, cache_policy=cp_i64, **kwargs)+ except (TypeError, ValueError):+ # Older DSL / op variant without cache_policy kw, or wrong scalar type+ pass+ return cute.copy(atom, src, dst, **kwargs)++class Sm100BlockScaledPersistentDenseGemmKernel:"""Persistent block-scaled GEMM kernel for SM100 (Blackwell).C = A x SFA x B x SFB (block-scaled NVF4 path)Notes:- A/B: Float4E2M1FN- - SF : Float8E4M3FN+ - SF : Float8E4M3FN (vector size 16)- Acc: Float32- C : Float16/BFloat16/Float32+ - A/B are K-major in this challenge setting."""def __init__(⋯ 695 unchanged linesab_producer_state)- for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):+ for _k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)- cute.copy(+ _cute_copy_with_cache_policy(tma_atom_a,tAgA_slice[(None, ab_producer_state.count)],tAsA[(None, ab_producer_state.index)],+ cache_policy=_CACHE_A,tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=a_full_mcast_mask,)- cute.copy(+ _cute_copy_with_cache_policy(tma_atom_b,tBgB_slice[(None, ab_producer_state.count)],tBsB[(None, ab_producer_state.index)],+ cache_policy=_CACHE_B,tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=b_full_mcast_mask,)- cute.copy(+ _cute_copy_with_cache_policy(tma_atom_sfa,tAgSFA_slice[(None, ab_producer_state.count)],tAsSFA[(None, ab_producer_state.index)],+ cache_policy=_CACHE_SFA,tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=sfa_full_mcast_mask,)- cute.copy(+ _cute_copy_with_cache_policy(tma_atom_sfb,tBgSFB_slice[(None, ab_producer_state.count)],tBsSFB[(None, ab_producer_state.index)],+ cache_policy=_CACHE_SFB,tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),mcast_mask=sfb_full_mcast_mask,)⋯ 116 unchanged lines# Reset ACCUMULATE each output tiletiled_mma.set(tcgen05.Field.ACCUMULATE, False)- # Patch B: hoist num_kblocks+ # Hoist num_kblocksnum_kblocks = cute.size(tCrA, mode=[2])for k_tile in range(k_tile_cnt):⋯ 47 unchanged linestCtAcc,)- # Patch C: set ACCUMULATE=True exactly once per output tile+ # Set ACCUMULATE=True exactly once per output tileif k_tile == 0:if kblock_idx == 0:tiled_mma.set(tcgen05.Field.ACCUMULATE, True)⋯ 71 unchanged linescur_tile_coord[2],)- bSG_gC = bSG_gC_partitioned[- (None, None, None, *mma_tile_coord_mnl)- ]+ bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_consumer_state.index)]⋯ 27 unchanged linesself.epilog_sync_barrier.arrive_and_wait()if warp_idx == self.epilog_warp_id[0]:- cute.copy(+ _cute_copy_with_cache_policy(tma_atom_c,bSG_sC[(None, c_buffer)],bSG_gC[(None, subtile_idx)],+ cache_policy=_CACHE_C,)c_pipeline.producer_commit()c_pipeline.producer_acquire()⋯ 188 unchanged lines)return tile_sched_params, grid- @staticmethod- def is_valid_dtypes_and_scale_factor_vec_size(- ab_dtype: Type[cutlass.Numeric],- sf_dtype: Type[cutlass.Numeric],- sf_vec_size: int,- c_dtype: Type[cutlass.Numeric],- ) -> bool:- is_valid = True- if ab_dtype not in {cutlass.Float4E2M1FN}:- is_valid = False- if sf_vec_size not in {16}:- is_valid = False- if sf_dtype not in {cutlass.Float8E4M3FN}:- is_valid = False- if c_dtype not in {cutlass.Float32, cutlass.Float16, cutlass.BFloat16}:- is_valid = False- return is_valid- @staticmethod- def is_valid_layouts(- ab_dtype: Type[cutlass.Numeric],- c_dtype: Type[cutlass.Numeric],- a_major: str,- b_major: str,- c_major: str,- ) -> bool:- is_valid = True- if ab_dtype is cutlass.Float4E2M1FN and not (a_major == "k" and b_major == "k"):- is_valid = False- if ab_dtype is cutlass.Float4E2M1FN and c_major == "m":- is_valid = False- return is_valid-- @staticmethod- def is_valid_mma_tiler_and_cluster_shape(- mma_tiler_mn: Tuple[int, int],- cluster_shape_mn: Tuple[int, int],- ) -> bool:- is_valid = True- if mma_tiler_mn[0] not in [128, 256]:- is_valid = False- if mma_tiler_mn[1] not in [64, 128, 192, 256]:- is_valid = False- if cluster_shape_mn[0] % (2 if mma_tiler_mn[0] == 256 else 1) != 0:- is_valid = False- is_power_of_2 = lambda x: x > 0 and (x & (x - 1)) == 0- if (- cluster_shape_mn[0] * cluster_shape_mn[1] > 16- or cluster_shape_mn[0] <= 0- or cluster_shape_mn[1] <= 0- or cluster_shape_mn[0] > 4- or cluster_shape_mn[1] > 4- or not is_power_of_2(cluster_shape_mn[0])- or not is_power_of_2(cluster_shape_mn[1])- ):- is_valid = False- return is_valid-- @staticmethod- def is_valid_tensor_alignment(- m: int,- n: int,- k: int,- l: int,- ab_dtype: Type[cutlass.Numeric],- c_dtype: Type[cutlass.Numeric],- a_major: str,- b_major: str,- c_major: str,- ) -> bool:- is_valid = True-- def check_contigous_16B_alignment(dtype, is_mode0_major, tensor_shape):- major_mode_idx = 0 if is_mode0_major else 1- num_major_elements = tensor_shape[major_mode_idx]- num_contiguous_elements = 16 * 8 // dtype.width- return num_major_elements % num_contiguous_elements == 0-- if (- not check_contigous_16B_alignment(ab_dtype, a_major == "m", (m, k, l))- or not check_contigous_16B_alignment(ab_dtype, b_major == "n", (n, k, l))- or not check_contigous_16B_alignment(c_dtype, c_major == "m", (m, n, l))- ):- is_valid = False- return is_valid-- @staticmethod- def can_implement(- ab_dtype: Type[cutlass.Numeric],- sf_dtype: Type[cutlass.Numeric],- sf_vec_size: int,- c_dtype: Type[cutlass.Numeric],- mma_tiler_mn: Tuple[int, int],- cluster_shape_mn: Tuple[int, int],- m: int,- n: int,- k: int,- l: int,- a_major: str,- b_major: str,- c_major: str,- ) -> bool:- can_implement = True- if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_dtypes_and_scale_factor_vec_size(- ab_dtype, sf_dtype, sf_vec_size, c_dtype- ):- can_implement = False- if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_layouts(- ab_dtype, c_dtype, a_major, b_major, c_major- ):- can_implement = False- if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_mma_tiler_and_cluster_shape(- mma_tiler_mn, cluster_shape_mn- ):- can_implement = False- if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_tensor_alignment(- m, n, k, l, ab_dtype, c_dtype, a_major, b_major, c_major- ):- can_implement = False- return can_implement--# ----------------- Hackathon wrapper -----------------# Data typesab_dtype = cutlass.Float4E2M1FNsf_dtype = cutlass.Float8E4M3FN- c_dtype = cutlass.Float16+ c_dtype = cutlass.Float16# Default configuration (baseline)_SF_VEC_SIZE = 16⋯ 3 unchanged lines# Compile options_COMPILE_OPT_LEVEL = 3- # Single supported K-tiling (IMPORTANT: avoid useless compile-fail + fallback noise)+ # Single supported K-tiling (avoid compile-fail + fallback noise)_GEMM = Sm100BlockScaledPersistentDenseGemmKernel(sf_vec_size=_SF_VEC_SIZE,mma_tiler_mn=_MMA_TILER_MN,⋯ 2 unchanged lines)# Compile cache keyed by problem_size (m,n,k,l)- _compiled = {}- _max_active_clusters = {}+ _compiled: Dict[Tuple[int, int, int, int], object] = {}+ # NOTE: Do NOT query HardwareInfo at import time.+ #+ # The benchmark harness imports this module inside multiprocessing workers.+ # At import time there may be no active CUDA context, and `utils.HardwareInfo()`+ # uses CUDA driver APIs that require a valid current context. That can trigger:+ # CUDA_ERROR_INVALID_CONTEXT+ #+ # For our fixed cluster shape (1,1) and the three target problems, the total+ # number of output tiles is <= 112, so as long as max_active_clusters is >= 112+ # it will not cap the grid. We therefore use a conservative constant to avoid+ # any driver calls during import.+ _MAX_ACTIVE_CLUSTERS = 1024- def _get_max_active_clusters(cluster_shape_mn):- cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]- if cluster_size in _max_active_clusters:- return _max_active_clusters[cluster_size]- hw = utils.HardwareInfo()- val = hw.get_max_active_clusters(cluster_size)- _max_active_clusters[cluster_size] = val- return val-- def compile_kernel(problem_size):- """Compile and cache a kernel specialized for (m,n,k,l).-- NOTE:- - We intentionally DO NOT try experimental ktile variants here.- - Some ktile variants fail to compile due to SFA TMA layout/vmap mismatch,- producing noisy stderr and wasting compile time.- """+ def compile_kernel(problem_size: Tuple[int, int, int, int]):+ """Compile and cache a kernel specialized for (m,n,k,l)."""ps = tuple(problem_size)if ps in _compiled:return _compiled[ps]⋯ 5 unchanged linessfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)+ # A/B are K-major for this task; output is row-major.layouts = (tcgen05.OperandMajorMode.K,tcgen05.OperandMajorMode.K,utils.LayoutEnum.ROW_MAJOR,)- max_active_clusters = _get_max_active_clusters(_CLUSTER_SHAPE_MN)-compiled = cute.compile(_GEMM,- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,+ a_ptr,+ b_ptr,+ sfa_ptr,+ sfb_ptr,+ c_ptr,layouts,ps,- max_active_clusters,+ _MAX_ACTIVE_CLUSTERS,options=f"--opt-level {_COMPILE_OPT_LEVEL}",)⋯ 15 unchanged linesa_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- sfa_ptr = make_ptr(sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- sfb_ptr = make_ptr(sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)+ sfa_ptr = make_ptr(+ sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16+ )+ sfb_ptr = make_ptr(+ sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=16+ )c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- # The compiled callable uses the default CUDA execution queue.+ # Constexpr args (layouts/max_active_clusters) are baked in by cute.compile.compiled(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, problem_size)return c
scrolls · 431 diff lines total
Best evidence level for this revision: reported
JSON