submission 203671
Simon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1424 lines, June 9 Researcher Reciprocity License v1.0.
better_baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-203671?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:21fef3fcd602cc3c047098493ed703a7dfcaf52571eee88867cf43cd6b87d2a6
license declaredunknown
license concludedunknown
authorsSimon
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
tmem_alloc_barrier = pipeline.NamedBarrier(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
better_baseline.py1424 lines
from typing import Type, Tuple, Union
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
class Sm100BlockScaledDenseGemmKernel:
def __init__(
self,
sf_vec_size: int,
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
):
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)
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.occupancy = 1
self.threads_per_cta = 512
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
SM100_TMEM_CAPACITY_COLUMNS = 512
self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS
def _setup_attributes(self):
"""Set up configurations that are dependent on GEMM inputs"""
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 = 4
self.mma_tiler = (
self.mma_inst_shape_mn[0],
self.mma_inst_shape_mn[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.mma_tiler_sfb = (
self.mma_inst_shape_mn_sfb[0],
self.mma_inst_shape_mn_sfb[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.cta_tile_shape_mnk = (
self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler[1],
self.mma_tiler[2],
)
self.cta_tile_shape_mnk_sfb = (
self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler_sfb[1],
self.mma_tiler_sfb[2],
)
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,),
)
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
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])
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,
)
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,
)
sf_atom_mn = 32
mma_inst_tile_k = 4
self.num_sfa_tmem_cols = (
self.cta_tile_shape_mnk[0] // sf_atom_mn
) * mma_inst_tile_k
self.num_sfb_tmem_cols = (
self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn
) * mma_inst_tile_k
self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage
@cute.jit
def __call__(
self,
a_ptr: cute.Pointer,
b1_ptr: cute.Pointer,
b2_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb1_ptr: cute.Pointer,
sfb2_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: cutlass.Constexpr,
epilogue_op: cutlass.Constexpr = lambda x: x
* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
"""Execute the dual GEMM operation: C = silu(A @ B1) * (A @ B2)"""
m, n, k, l = problem_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, k, l),
stride=(k, 1, m * k),
),
)
b_tensor1 = cute.make_tensor(
b1_ptr,
cute.make_layout(
(n, k, l),
stride=(k, 1, n * k),
),
)
b_tensor2 = cute.make_tensor(
b2_ptr,
cute.make_layout(
(n, k, l),
stride=(k, 1, n * k),
),
)
c_tensor = cute.make_tensor(
c_ptr, cute.make_layout((m, n, l), stride=(n, 1, m * n))
)
self.a_dtype: Type[cutlass.Numeric] = a_tensor.element_type
self.b_dtype: Type[cutlass.Numeric] = b_tensor1.element_type
self.sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN
self.c_dtype: Type[cutlass.Numeric] = c_tensor.element_type
self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor1).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(c_tensor)
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()
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_tensor1.shape, self.sf_vec_size
)
sfb_tensor1 = cute.make_tensor(sfb1_ptr, sfb_layout)
sfb_tensor2 = cute.make_tensor(sfb2_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)
# Setup TMA load for 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,
)
# Setup TMA load for B1 and B2
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_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor1,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor2,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
# Setup TMA load for 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,
)
# Setup TMA load for SFB1 and SFB2
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_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor1,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor2,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x = tma_tensor_sfb1.stride[0][1]
y = cute.ceil_div(tma_tensor_sfb1.shape[0][1], 4)
new_shape = (
(tma_tensor_sfb1.shape[0][0], ((2, 2), y)),
tma_tensor_sfb1.shape[1],
tma_tensor_sfb1.shape[2],
)
x_times_3 = 3 * x
new_stride = (
(tma_tensor_sfb1.stride[0][0], ((x, x), x_times_3)),
tma_tensor_sfb1.stride[1],
tma_tensor_sfb1.stride[2],
)
tma_tensor_sfb_new_layout = cute.make_layout(new_shape, stride=new_stride)
tma_tensor_sfb1 = cute.make_tensor(
tma_tensor_sfb1.iterator, tma_tensor_sfb_new_layout
)
tma_tensor_sfb2 = cute.make_tensor(
tma_tensor_sfb2.iterator, tma_tensor_sfb_new_layout
)
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 * 2 + sfa_copy_size + sfb_copy_size * 2
) * atom_thr_size
# Setup 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,
)
grid = self._compute_grid(
c_tensor,
self.cta_tile_shape_mnk,
self.cluster_shape_mn,
)
self.buffer_align_bytes = 1024
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2]
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
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,
]
sB1: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
sB2: 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,
]
sSFB1: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
],
self.buffer_align_bytes,
]
sSFB2: 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_b1,
tma_tensor_b1,
tma_atom_b2,
tma_tensor_b2,
tma_atom_sfa,
tma_tensor_sfa,
tma_atom_sfb1,
tma_tensor_sfb1,
tma_atom_sfb2,
tma_tensor_sfb2,
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,
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_b1: cute.CopyAtom,
mB_nkl1: cute.Tensor,
tma_atom_b2: cute.CopyAtom,
mB_nkl2: cute.Tensor,
tma_atom_sfa: cute.CopyAtom,
mSFA_mkl: cute.Tensor,
tma_atom_sfb1: cute.CopyAtom,
mSFB_nkl1: cute.Tensor,
tma_atom_sfb2: cute.CopyAtom,
mSFB_nkl2: 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,
epilogue_op: cutlass.Constexpr,
):
"""GPU device kernel performing the batched dual GEMM computation with block scaling."""
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# Prefetch tma desc
if warp_idx == 0:
cpasync.prefetch_descriptor(tma_atom_a)
cpasync.prefetch_descriptor(tma_atom_b1)
cpasync.prefetch_descriptor(tma_atom_b2)
cpasync.prefetch_descriptor(tma_atom_sfa)
cpasync.prefetch_descriptor(tma_atom_sfb1)
cpasync.prefetch_descriptor(tma_atom_sfb2)
cpasync.prefetch_descriptor(tma_atom_c)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
# Setup 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
)
cta_coord = (bidx, bidy, bidz)
mma_tile_coord_mnl = (
cta_coord[0] // cute.size(tiled_mma.thr_id.shape),
cta_coord[1],
cta_coord[2],
)
tidx, _, _ = cute.arch.thread_idx()
# Alloc and init barriers
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
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_producer, ab_consumer = 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,
).make_participants()
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, self.threads_per_cta
)
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,
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=2, num_threads=self.threads_per_cta
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=tmem_alloc_barrier,
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
)
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
)
sB1 = storage.sB1.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sB2 = storage.sB2.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB1 = storage.sSFB1.get_tensor(sfb_smem_layout_staged)
sSFB2 = storage.sSFB2.get_tensor(sfb_smem_layout_staged)
# Compute 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_tile partition global tensors
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gB_nkl1 = cute.local_tile(
mB_nkl1, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
gB_nkl2 = cute.local_tile(
mB_nkl2, 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_nkl1 = cute.local_tile(
mSFB_nkl1,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
gSFB_nkl2 = cute.local_tile(
mSFB_nkl2,
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 global tensor 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)
tCgB1 = thr_mma.partition_B(gB_nkl1)
tCgB2 = thr_mma.partition_B(gB_nkl2)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB1 = thr_mma_sfb.partition_B(gSFB_nkl1)
tCgSFB2 = thr_mma_sfb.partition_B(gSFB_nkl2)
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 B1 and B2
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB1, tBgB1 = cpasync.tma_partition(
tma_atom_b1,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB1, 0, 3),
cute.group_modes(tCgB1, 0, 3),
)
tBsB2, tBgB2 = cpasync.tma_partition(
tma_atom_b2,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB2, 0, 3),
cute.group_modes(tCgB2, 0, 3),
)
# TMA partitions for SFA
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)
# TMA partitions for SFB1 and SFB2
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB1, tBgSFB1 = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb1,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB1, 0, 3),
cute.group_modes(tCgSFB1, 0, 3),
)
tBsSFB1 = cute.filter_zeros(tBsSFB1)
tBgSFB1 = cute.filter_zeros(tBgSFB1)
tBsSFB2, tBgSFB2 = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb2,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB2, 0, 3),
cute.group_modes(tCgSFB2, 0, 3),
)
tBsSFB2 = cute.filter_zeros(tBsSFB2)
tBgSFB2 = cute.filter_zeros(tBgSFB2)
# Partition shared/tensor memory for TiledMMA
tCrA = tiled_mma.make_fragment_A(sA)
tCrB1 = tiled_mma.make_fragment_B(sB1)
tCrB2 = tiled_mma.make_fragment_B(sB2)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
tCtAcc_fake = tiled_mma.make_fragment_C(acc_shape)
# Cluster wait before tensor memory alloc
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
# Create two accumulators in tensor memory
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc1 = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
acc_tmem_ptr1 = cute.recast_ptr(
acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),
dtype=self.acc_dtype,
)
tCtAcc2 = cute.make_tensor(acc_tmem_ptr1, tCtAcc_fake.layout)
# Make SFA tmem tensor
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols * 2,
dtype=self.sf_dtype,
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
# Make SFB1 and SFB2 tmem tensors
sfb_tmem_ptr1 = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols * 2 + 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)),
)
tCtSFB1 = cute.make_tensor(sfb_tmem_ptr1, tCtSFB_layout)
sfb_tmem_ptr2 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols * 2
+ self.num_sfa_tmem_cols
+ self.num_sfb_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB2 = cute.make_tensor(sfb_tmem_ptr2, tCtSFB_layout)
# Partition for S2T copy of SFA/SFB
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB1_compact_s2t,
tCtSFB1_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB1, tCtSFB1)
# Partition SFB2 using the same tiled copy
tCsSFB2_compact = cute.filter_zeros(sSFB2)
tCtSFB2_compact = cute.filter_zeros(tCtSFB2)
thr_copy_s2t_sfb = tiled_copy_s2t_sfb.get_slice(0)
tCsSFB2_compact_s2t_ = thr_copy_s2t_sfb.partition_S(tCsSFB2_compact)
tCsSFB2_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t_sfb, tCsSFB2_compact_s2t_
)
tCtSFB2_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB2_compact)
# Slice to per mma tile index
tAgA = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB1 = tBgB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tBgB2 = tBgB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA = 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
tBgSFB1 = tBgSFB1[(None, slice_n, None, mma_tile_coord_mnl[2])]
tBgSFB2 = tBgSFB2[(None, slice_n, None, mma_tile_coord_mnl[2])]
tCtSFB1_mma = tCtSFB1
tCtSFB2_mma = tCtSFB2
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_ptr1 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols * 2
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB1_mma = cute.make_tensor(shifted_ptr1, tCtSFB_layout)
shifted_ptr2 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols * 2
+ self.num_sfa_tmem_cols
+ self.num_sfb_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB2_mma = cute.make_tensor(shifted_ptr2, tCtSFB_layout)
elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
shifted_ptr1 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols * 2
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB1_mma = cute.make_tensor(shifted_ptr1, tCtSFB_layout)
shifted_ptr2 = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols * 2
+ self.num_sfa_tmem_cols
+ self.num_sfb_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB2_mma = cute.make_tensor(shifted_ptr2, tCtSFB_layout)
# Pipelining TMA load and MMA mainloop
prefetch_k_tile_cnt = cutlass.min(self.num_ab_stage - 2, k_tile_cnt)
if warp_idx == 0:
# Prefetch TMA loads
for k_tile_idx in cutlass.range(prefetch_k_tile_cnt, unroll=1):
producer_handle = ab_producer.acquire_and_advance()
cute.copy(
tma_atom_a,
tAgA[(None, k_tile_idx)],
tAsA[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=a_full_mcast_mask,
)
cute.copy(
tma_atom_b1,
tBgB1[(None, k_tile_idx)],
tBsB1[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_b2,
tBgB2[(None, k_tile_idx)],
tBsB2[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_sfa,
tAgSFA[(None, k_tile_idx)],
tAsSFA[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfa_full_mcast_mask,
)
cute.copy(
tma_atom_sfb1,
tBgSFB1[(None, k_tile_idx)],
tBsSFB1[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfb_full_mcast_mask,
)
cute.copy(
tma_atom_sfb2,
tBgSFB2[(None, k_tile_idx)],
tBsSFB2[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfb_full_mcast_mask,
)
peek_ab_full_status = cutlass.Boolean(False)
if is_leader_cta:
peek_ab_full_status = ab_consumer.try_wait()
peek_ab_empty_status = ab_producer.try_acquire()
# MMA mainloop
for k_tile_idx in cutlass.range(k_tile_cnt):
if k_tile_idx < k_tile_cnt - prefetch_k_tile_cnt:
producer_handle = ab_producer.acquire_and_advance(
peek_ab_empty_status
)
cute.copy(
tma_atom_a,
tAgA[(None, producer_handle.count)],
tAsA[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=a_full_mcast_mask,
)
cute.copy(
tma_atom_b1,
tBgB1[(None, producer_handle.count)],
tBsB1[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_b2,
tBgB2[(None, producer_handle.count)],
tBsB2[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_sfa,
tAgSFA[(None, producer_handle.count)],
tAsSFA[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfa_full_mcast_mask,
)
cute.copy(
tma_atom_sfb1,
tBgSFB1[(None, producer_handle.count)],
tBsSFB1[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfb_full_mcast_mask,
)
cute.copy(
tma_atom_sfb2,
tBgSFB2[(None, producer_handle.count)],
tBsSFB2[(None, producer_handle.index)],
tma_bar_ptr=producer_handle.barrier,
mcast_mask=sfb_full_mcast_mask,
)
if is_leader_cta:
consumer_handle = ab_consumer.wait_and_advance(peek_ab_full_status)
# Copy SFA/SFB from smem to tmem
s2t_stage_coord = (
None,
None,
None,
None,
consumer_handle.index,
)
tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
tCsSFB1_compact_s2t_staged = tCsSFB1_compact_s2t[s2t_stage_coord]
tCsSFB2_compact_s2t_staged = tCsSFB2_compact_s2t[s2t_stage_coord]
cute.copy(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t_staged,
tCtSFA_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB1_compact_s2t_staged,
tCtSFB1_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB2_compact_s2t_staged,
tCtSFB2_compact_s2t,
)
# Dual GEMM computation
num_kblks = cute.size(tCrA, mode=[2])
for kblk_idx in cutlass.range(num_kblks, unroll_full=True):
kblk_crd = (None, None, kblk_idx, consumer_handle.index)
sf_kblock_coord = (None, None, kblk_idx)
# First GEMM: Acc1 = A @ B1
tiled_mma.set(
tcgen05.Field.SFA,
tCtSFA[sf_kblock_coord].iterator,
)
tiled_mma.set(
tcgen05.Field.SFB,
tCtSFB1_mma[sf_kblock_coord].iterator,
)
cute.gemm(
tiled_mma,
tCtAcc1,
tCrA[kblk_crd],
tCrB1[kblk_crd],
tCtAcc1,
)
# Second GEMM: Acc2 = A @ B2
tiled_mma.set(
tcgen05.Field.SFB,
tCtSFB2_mma[sf_kblock_coord].iterator,
)
cute.gemm(
tiled_mma,
tCtAcc2,
tCrA[kblk_crd],
tCrB2[kblk_crd],
tCtAcc2,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
consumer_handle.release()
peek_ab_empty_status = ab_producer.try_acquire()
peek_ab_full_status = ab_consumer.try_wait()
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
# Epilogue
tmem.relinquish_alloc_permit()
acc_pipeline.consumer_wait(acc_consumer_state)
self.epilogue_tma_store(
tidx,
warp_idx,
mma_tile_coord_mnl,
tma_atom_c,
tCtAcc1,
tCtAcc2,
sC,
tCgC,
epi_tile,
epilogue_op,
)
pipeline.sync(barrier_id=1)
tmem.free(acc_tmem_ptr)
if warp_idx == 0:
ab_producer.tail()
def mainloop_s2t_copy_and_partition(
self,
sSF: cute.Tensor,
tSF: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
"""Make tiledCopy for smem to tmem load for scale factor 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]:
"""Make tiledCopy for tensor memory load."""
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0)], epi_tile)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 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]:
"""Make tiledCopy for shared memory store."""
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
@cute.jit
def epilogue_tma_store(
self,
epi_tidx: cutlass.Int32,
warp_idx: cutlass.Int32,
mma_tile_coord_mnl: Tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32],
tma_atom_c: cute.CopyAtom,
tCtAcc1: cute.Tensor,
tCtAcc2: cute.Tensor,
sC: cute.Tensor,
tCgC: cute.Tensor,
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
) -> None:
"""Epilogue implementation for TMA store version with dual accumulator fusion."""
tiled_copy_t2r, tTR_tAcc1, tTR_rAcc1 = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc1, tCgC, epi_tile, self.use_2cta_instrs
)
tTR_tAcc1 = cute.group_modes(tTR_tAcc1, 3, cute.rank(tTR_tAcc1))
# Partition for second accumulator
tAcc2_epi = cute.flat_divide(tCtAcc2[((None, None), 0, 0)], epi_tile)
thr_copy_t2r = tiled_copy_t2r.get_slice(epi_tidx)
tTR_tAcc2 = thr_copy_t2r.partition_S(tAcc2_epi)
tTR_tAcc2 = cute.group_modes(tTR_tAcc2, 3, cute.rank(tTR_tAcc2))
tTR_rAcc2 = cute.make_rmem_tensor(tTR_rAcc1.shape, self.acc_dtype)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc1.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
)
tCgC_epi = cute.flat_divide(
tCgC[((None, None), 0, 0, None, None, None)], epi_tile
)
bSG_sC, bSG_gC = cpasync.tma_partition(
tma_atom_c,
0,
cute.make_layout(1),
cute.group_modes(sC, 0, 2),
cute.group_modes(tCgC_epi, 0, 2),
)
bSG_gC = bSG_gC[(None, None, None, *mma_tile_coord_mnl)]
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, self.threads_per_cta
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage, producer_group=c_producer_group
)
subtile_cnt = cute.size(tTR_tAcc1.shape, mode=[3])
for subtile_idx in cutlass.range(subtile_cnt):
tTR_tAcc1_mn = tTR_tAcc1[(None, None, None, subtile_idx)]
tTR_tAcc2_mn = tTR_tAcc2[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc1_mn, tTR_rAcc1)
cute.copy(tiled_copy_t2r, tTR_tAcc2_mn, tTR_rAcc2)
# Apply silu to acc1 and multiply with acc2: silu(acc1) * acc2
acc_vec1 = tiled_copy_r2s.retile(tTR_rAcc1).load()
acc_vec2 = tiled_copy_r2s.retile(tTR_rAcc2).load()
acc_vec1 = epilogue_op(acc_vec1) # silu(acc1)
acc_vec = acc_vec1 * acc_vec2
tRS_rC.store(acc_vec.to(self.c_dtype))
c_buffer = 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,
)
pipeline.sync(barrier_id=1)
if warp_idx == 0:
cute.copy(
tma_atom_c, bSG_sC[(None, c_buffer)], bSG_gC[(None, subtile_idx)]
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
pipeline.sync(barrier_id=1)
c_pipeline.producer_tail()
@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]:
"""Computes the number of stages for A/B/C operands based on heuristics."""
num_acc_stage = 1
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,
)
# Account for dual B and SFB buffers
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) * 2 # B1 and B2
+ cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
* 2 # SFB1 and SFB2
)
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 + 1) * (mbar_helpers_bytes + c_bytes)
) // ab_bytes_per_stage
num_c_stage += (
smem_capacity
- ab_bytes_per_stage * num_ab_stage
- (occupancy + 1) * (mbar_helpers_bytes + c_bytes)
) // ((occupancy + 1) * 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],
) -> Tuple[int, int, int]:
"""Compute grid shape for the output tensor C."""
cluster_shape_mnl = (*cluster_shape_mn, 1)
grid = cute.round_up(
(
cute.ceil_div(c.layout.shape[0], cta_tile_shape_mnk[0]),
cute.ceil_div(c.layout.shape[1], cta_tile_shape_mnk[1]),
c.layout.shape[2],
),
cluster_shape_mnl,
)
return grid
# FP4 data type for A and B
ab_dtype = cutlass.Float4E2M1FN
# FP8 data type for scale factors
sf_dtype = cutlass.Float8E4M3FN
# FP16 output type
c_dtype = cutlass.Float16
# Scale factor block size (16 elements share one scale)
sf_vec_size = 16
# Global cache for compiled kernel
_compiled_kernel_cache = {}
def compile_kernel(problem_size):
"""
Compile the kernel once and cache it for each problem size.
This should be called before any timing measurements.
Returns:
The compiled kernel function
"""
global _compiled_kernel_cache
if problem_size in _compiled_kernel_cache:
return _compiled_kernel_cache[problem_size]
# Create CuTe pointers for A/B1/B2/C/SFA/SFB1/SFB2
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b1_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b2_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)
sfb1_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb2_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
mma_tiler_mn = (128, 64)
cluster_shape_mn = (1, 2)
gemm = Sm100BlockScaledDenseGemmKernel(
sf_vec_size,
mma_tiler_mn,
cluster_shape_mn,
)
_compiled_kernel_cache[problem_size] = cute.compile(
gemm,
a_ptr,
b1_ptr,
b2_ptr,
sfa_ptr,
sfb1_ptr,
sfb2_ptr,
c_ptr,
problem_size,
options="--opt-level 3",
)
return _compiled_kernel_cache[problem_size]
def custom_kernel(data):
"""
Execute the block-scaled dual GEMM kernel with silu activation.
C = silu(A @ B1) * (A @ B2)
Args:
data: Tuple of (a, b1, b2, sfa_ref, sfb1_ref, sfb2_ref, sfa_permuted, sfb1_permuted, sfb2_permuted, c) PyTorch tensors
Returns:
Output tensor c with computed results
"""
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
# Get dimensions from MxKxL layout
m, k, l = a.shape
n, _, _ = b1.shape
# Torch uses e2m1_x2 data type, thus k is halved
k = k * 2
# Compile kernel for this specific problem size
problem_size = (m, n, k, l)
compiled_func = compile_kernel(problem_size)
# Create CuTe pointers
a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b1_ptr = make_ptr(ab_dtype, b1.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b2_ptr = make_ptr(ab_dtype, b2.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(
sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
sfb1_ptr = make_ptr(
sf_dtype, sfb1_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
sfb2_ptr = make_ptr(
sf_dtype, sfb2_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
# Execute the compiled kernel (problem_size is baked in at compile time)
compiled_func(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr)
return c
scrolls · 1424 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 202766.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON