submission 202766
Simon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 907 lines, June 9 Researcher Reciprocity License v1.0.
baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-202766?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:25cf04a32ed4874ec70c97db6c952c3fa7d4d951ee9f8ebef589954ef6666725
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
epilogue_op: cutlass.Constexpr = lambda x: xmbarrier
tmem_alloc_barrier = pipeline.NamedBarrier(shared-memory
a_smem_layout_staged: cute.ComposedLayout,tcgen05
acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1),warp-specialization
ab1_producer, ab1_consumer = pipeline.PipelineTmaUmma.create(Kernel source
baseline.py907 lines
from task import input_t, output_t
import cutlass
import cutlass.cute as cute
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import make_ptr
# Kernel configuration parameters
mma_tiler_mnk = (128, 128, 256)
mma_inst_shape_k = 64
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
sf_vec_size = 16
threads_per_cta = 128 # 4 warps
# Separate stages for each pipeline
num_ab1_stage = 3
num_ab2_stage = 3
num_acc_stage = 1
num_tmem_alloc_cols = 512
@cute.kernel
def kernel(
tiled_mma: cute.TiledMma,
# GEMM1: A1, B1, SFA1, SFB1
tma_atom_a1: cute.CopyAtom,
mA_mkl1: cute.Tensor,
tma_atom_b1: cute.CopyAtom,
mB_nkl1: cute.Tensor,
tma_atom_sfa1: cute.CopyAtom,
mSFA_mkl1: cute.Tensor,
tma_atom_sfb1: cute.CopyAtom,
mSFB_nkl1: cute.Tensor,
# GEMM2: A2, B2, SFA2, SFB2
tma_atom_a2: cute.CopyAtom,
mA_mkl2: cute.Tensor,
tma_atom_b2: cute.CopyAtom,
mB_nkl2: cute.Tensor,
tma_atom_sfa2: cute.CopyAtom,
mSFA_mkl2: cute.Tensor,
tma_atom_sfb2: cute.CopyAtom,
mSFB_nkl2: cute.Tensor,
# Output
mC_mnl: cute.Tensor,
# Layouts
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
# Separate byte counts for each pipeline
num_tma_bytes_gemm1: cutlass.Constexpr[int],
num_tma_bytes_gemm2: cutlass.Constexpr[int],
epilogue_op: cutlass.Constexpr = lambda x: x
* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
tidx, _, _ = cute.arch.thread_idx()
bidx, bidy, bidz = cute.arch.block_idx()
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
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],
)
#
# Shared storage with separate barriers for each pipeline
#
@cute.struct
class SharedStorage:
ab1_bar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab1_stage * 2]
ab2_bar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab2_stage * 2]
acc1_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage * 2]
acc2_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage * 2]
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
storage = smem.allocate(SharedStorage)
# GEMM1 SMEM: A1, B1, SFA1, SFB1
sA1 = smem.allocate_tensor(
element_type=ab_dtype,
layout=a_smem_layout_staged.outer,
byte_alignment=128,
swizzle=a_smem_layout_staged.inner,
)
sB1 = smem.allocate_tensor(
element_type=ab_dtype,
layout=b_smem_layout_staged.outer,
byte_alignment=128,
swizzle=b_smem_layout_staged.inner,
)
sSFA1 = smem.allocate_tensor(
element_type=sf_dtype,
layout=sfa_smem_layout_staged,
byte_alignment=128,
)
sSFB1 = smem.allocate_tensor(
element_type=sf_dtype,
layout=sfb_smem_layout_staged,
byte_alignment=128,
)
# GEMM2 SMEM: A2, B2, SFA2, SFB2
sA2 = smem.allocate_tensor(
element_type=ab_dtype,
layout=a_smem_layout_staged.outer,
byte_alignment=128,
swizzle=a_smem_layout_staged.inner,
)
sB2 = smem.allocate_tensor(
element_type=ab_dtype,
layout=b_smem_layout_staged.outer,
byte_alignment=128,
swizzle=b_smem_layout_staged.inner,
)
sSFA2 = smem.allocate_tensor(
element_type=sf_dtype,
layout=sfa_smem_layout_staged,
byte_alignment=128,
)
sSFB2 = smem.allocate_tensor(
element_type=sf_dtype,
layout=sfb_smem_layout_staged,
byte_alignment=128,
)
#
# Four independent pipelines
#
# GEMM1 TMA pipeline: Warp 0 produces, Warp 2 consumes
ab1_producer, ab1_consumer = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab1_bar_ptr.data_ptr(),
num_stages=num_ab1_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1),
tx_count=num_tma_bytes_gemm1,
).make_participants()
# GEMM2 TMA pipeline: Warp 1 produces, Warp 3 consumes
ab2_producer, ab2_consumer = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab2_bar_ptr.data_ptr(),
num_stages=num_ab2_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1),
tx_count=num_tma_bytes_gemm2,
).make_participants()
# ACC1 pipeline: Warp 2 produces, All threads consume (for epilogue)
acc1_producer, acc1_consumer = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc1_mbar_ptr.data_ptr(),
num_stages=num_acc_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(
pipeline.Agent.Thread, threads_per_cta
),
).make_participants()
# ACC2 pipeline: Warp 3 produces, All threads consume (for epilogue)
acc2_producer, acc2_consumer = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc2_mbar_ptr.data_ptr(),
num_stages=num_acc_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(
pipeline.Agent.Thread, threads_per_cta
),
).make_participants()
#
# Partition global tensors
#
# GEMM1
gA1_mkl = cute.local_tile(
mA_mkl1, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gB1_nkl = cute.local_tile(
mB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gSFA1_mkl = cute.local_tile(
mSFA_mkl1, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gSFB1_nkl = cute.local_tile(
mSFB_nkl1, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
# GEMM2
gA2_mkl = cute.local_tile(
mA_mkl2, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gB2_nkl = cute.local_tile(
mB_nkl2, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gSFA2_mkl = cute.local_tile(
mSFA_mkl2, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gSFB2_nkl = cute.local_tile(
mSFB_nkl2, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(gA1_mkl, mode=[3])
#
# MMA partitions
#
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
# GEMM1
tCgA1 = thr_mma.partition_A(gA1_mkl)
tCgB1 = thr_mma.partition_B(gB1_nkl)
tCgSFA1 = thr_mma.partition_A(gSFA1_mkl)
tCgSFB1 = thr_mma.partition_B(gSFB1_nkl)
# GEMM2
tCgA2 = thr_mma.partition_A(gA2_mkl)
tCgB2 = thr_mma.partition_B(gB2_nkl)
tCgSFA2 = thr_mma.partition_A(gSFA2_mkl)
tCgSFB2 = thr_mma.partition_B(gSFB2_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
#
# TMA partitions for GEMM1
#
tAsA1, tAgA1 = cpasync.tma_partition(
tma_atom_a1,
0,
cute.make_layout(1),
cute.group_modes(sA1, 0, 3),
cute.group_modes(tCgA1, 0, 3),
)
tBsB1, tBgB1 = cpasync.tma_partition(
tma_atom_b1,
0,
cute.make_layout(1),
cute.group_modes(sB1, 0, 3),
cute.group_modes(tCgB1, 0, 3),
)
tAsSFA1, tAgSFA1 = cpasync.tma_partition(
tma_atom_sfa1,
0,
cute.make_layout(1),
cute.group_modes(sSFA1, 0, 3),
cute.group_modes(tCgSFA1, 0, 3),
)
tAsSFA1 = cute.filter_zeros(tAsSFA1)
tAgSFA1 = cute.filter_zeros(tAgSFA1)
tBsSFB1, tBgSFB1 = cpasync.tma_partition(
tma_atom_sfb1,
0,
cute.make_layout(1),
cute.group_modes(sSFB1, 0, 3),
cute.group_modes(tCgSFB1, 0, 3),
)
tBsSFB1 = cute.filter_zeros(tBsSFB1)
tBgSFB1 = cute.filter_zeros(tBgSFB1)
#
# TMA partitions for GEMM2
#
tAsA2, tAgA2 = cpasync.tma_partition(
tma_atom_a2,
0,
cute.make_layout(1),
cute.group_modes(sA2, 0, 3),
cute.group_modes(tCgA2, 0, 3),
)
tBsB2, tBgB2 = cpasync.tma_partition(
tma_atom_b2,
0,
cute.make_layout(1),
cute.group_modes(sB2, 0, 3),
cute.group_modes(tCgB2, 0, 3),
)
tAsSFA2, tAgSFA2 = cpasync.tma_partition(
tma_atom_sfa2,
0,
cute.make_layout(1),
cute.group_modes(sSFA2, 0, 3),
cute.group_modes(tCgSFA2, 0, 3),
)
tAsSFA2 = cute.filter_zeros(tAsSFA2)
tAgSFA2 = cute.filter_zeros(tAgSFA2)
tBsSFB2, tBgSFB2 = cpasync.tma_partition(
tma_atom_sfb2,
0,
cute.make_layout(1),
cute.group_modes(sSFB2, 0, 3),
cute.group_modes(tCgSFB2, 0, 3),
)
tBsSFB2 = cute.filter_zeros(tBsSFB2)
tBgSFB2 = cute.filter_zeros(tBgSFB2)
#
# MMA fragments
#
tCrA1 = tiled_mma.make_fragment_A(sA1)
tCrB1 = tiled_mma.make_fragment_B(sB1)
tCrA2 = tiled_mma.make_fragment_A(sA2)
tCrB2 = tiled_mma.make_fragment_B(sB2)
acc_shape = tiled_mma.partition_shape_C(mma_tiler_mnk[:2])
tCtAcc_fake = tiled_mma.make_fragment_C(acc_shape)
#
# TMEM allocation
#
tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=1, num_threads=threads_per_cta
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf, barrier_for_retrieve=tmem_alloc_barrier
)
tmem.allocate(num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(cutlass.Float32)
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=cutlass.Float32,
)
tCtAcc2 = cute.make_tensor(acc_tmem_ptr1, tCtAcc_fake.layout)
# SFA/SFB TMEM for GEMM1
tCtSFA1_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
sfa1_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc2),
dtype=sf_dtype,
)
tCtSFA1 = cute.make_tensor(sfa1_tmem_ptr, tCtSFA1_layout)
tCtSFB1_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
sfb1_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA1),
dtype=sf_dtype,
)
tCtSFB1 = cute.make_tensor(sfb1_tmem_ptr, tCtSFB1_layout)
# SFA/SFB TMEM for GEMM2
sfa2_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA1)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFB1),
dtype=sf_dtype,
)
tCtSFA2 = cute.make_tensor(sfa2_tmem_ptr, tCtSFA1_layout)
sfb2_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc1)
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc2)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA1)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFB1)
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA2),
dtype=sf_dtype,
)
tCtSFB2 = cute.make_tensor(sfb2_tmem_ptr, tCtSFB1_layout)
#
# S2T copy setup for GEMM1
#
copy_atom_s2t = cute.make_copy_atom(
tcgen05.Cp4x32x128bOp(tcgen05.CtaGroup.ONE), sf_dtype
)
tCsSFA1_compact = cute.filter_zeros(sSFA1)
tCtSFA1_compact = cute.filter_zeros(tCtSFA1)
tiled_copy_s2t_sfa = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFA1_compact)
thr_copy_s2t_sfa = tiled_copy_s2t_sfa.get_slice(0)
tCsSFA1_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t_sfa, thr_copy_s2t_sfa.partition_S(tCsSFA1_compact)
)
tCtSFA1_compact_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA1_compact)
tCsSFB1_compact = cute.filter_zeros(sSFB1)
tCtSFB1_compact = cute.filter_zeros(tCtSFB1)
tiled_copy_s2t_sfb = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFB1_compact)
thr_copy_s2t_sfb = tiled_copy_s2t_sfb.get_slice(0)
tCsSFB1_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t_sfb, thr_copy_s2t_sfb.partition_S(tCsSFB1_compact)
)
tCtSFB1_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB1_compact)
#
# S2T copy setup for GEMM2
#
tCsSFA2_compact = cute.filter_zeros(sSFA2)
tCtSFA2_compact = cute.filter_zeros(tCtSFA2)
tCsSFA2_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t_sfa, thr_copy_s2t_sfa.partition_S(tCsSFA2_compact)
)
tCtSFA2_compact_s2t = thr_copy_s2t_sfa.partition_D(tCtSFA2_compact)
tCsSFB2_compact = cute.filter_zeros(sSFB2)
tCtSFB2_compact = cute.filter_zeros(tCtSFB2)
tCsSFB2_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t_sfb, thr_copy_s2t_sfb.partition_S(tCsSFB2_compact)
)
tCtSFB2_compact_s2t = thr_copy_s2t_sfb.partition_D(tCtSFB2_compact)
#
# Slice to per-tile coords
#
tAgA1 = tAgA1[(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])]
tAgSFA1 = tAgSFA1[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgSFB1 = tBgSFB1[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgA2 = tAgA2[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB2 = tBgB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA2 = tAgSFA2[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgSFB2 = tBgSFB2[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
#
# WARP 0: TMA Producer for GEMM1
#
if warp_idx == 0:
for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
ab1_empty = ab1_producer.acquire_and_advance()
cute.copy(
tma_atom_a1,
tAgA1[(None, ab1_empty.count)],
tAsA1[(None, ab1_empty.index)],
tma_bar_ptr=ab1_empty.barrier,
)
cute.copy(
tma_atom_b1,
tBgB1[(None, ab1_empty.count)],
tBsB1[(None, ab1_empty.index)],
tma_bar_ptr=ab1_empty.barrier,
)
cute.copy(
tma_atom_sfa1,
tAgSFA1[(None, ab1_empty.count)],
tAsSFA1[(None, ab1_empty.index)],
tma_bar_ptr=ab1_empty.barrier,
)
cute.copy(
tma_atom_sfb1,
tBgSFB1[(None, ab1_empty.count)],
tBsSFB1[(None, ab1_empty.index)],
tma_bar_ptr=ab1_empty.barrier,
)
#
# WARP 1: TMA Producer for GEMM2
#
if warp_idx == 1:
for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
ab2_empty = ab2_producer.acquire_and_advance()
cute.copy(
tma_atom_a2,
tAgA2[(None, ab2_empty.count)],
tAsA2[(None, ab2_empty.index)],
tma_bar_ptr=ab2_empty.barrier,
)
cute.copy(
tma_atom_b2,
tBgB2[(None, ab2_empty.count)],
tBsB2[(None, ab2_empty.index)],
tma_bar_ptr=ab2_empty.barrier,
)
cute.copy(
tma_atom_sfa2,
tAgSFA2[(None, ab2_empty.count)],
tAsSFA2[(None, ab2_empty.index)],
tma_bar_ptr=ab2_empty.barrier,
)
cute.copy(
tma_atom_sfb2,
tBgSFB2[(None, ab2_empty.count)],
tBsSFB2[(None, ab2_empty.index)],
tma_bar_ptr=ab2_empty.barrier,
)
#
# WARP 2: MMA Consumer for GEMM1
#
if warp_idx == 2:
acc1_empty = acc1_producer.acquire_and_advance()
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
ab1_full = ab1_consumer.wait_and_advance()
# S2T for SFA1, SFB1
s2t_coord1 = (None, None, None, None, ab1_full.index)
cute.copy(
tiled_copy_s2t_sfa, tCsSFA1_compact_s2t[s2t_coord1], tCtSFA1_compact_s2t
)
cute.copy(
tiled_copy_s2t_sfb, tCsSFB1_compact_s2t[s2t_coord1], tCtSFB1_compact_s2t
)
# MMA1: ACC1 += A1 @ B1
num_kblocks = cute.size(tCrA1, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, ab1_full.index)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(tcgen05.Field.SFA, tCtSFA1[sf_kblock_coord].iterator)
tiled_mma.set(tcgen05.Field.SFB, tCtSFB1[sf_kblock_coord].iterator)
cute.gemm(
tiled_mma,
tCtAcc1,
tCrA1[kblock_coord],
tCrB1[kblock_coord],
tCtAcc1,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab1_full.release()
acc1_empty.commit()
#
# WARP 3: MMA Consumer for GEMM2
#
if warp_idx == 3:
acc2_empty = acc2_producer.acquire_and_advance()
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in cutlass.range(k_tile_cnt, unroll_full=True):
ab2_full = ab2_consumer.wait_and_advance()
# S2T for SFA2, SFB2
s2t_coord2 = (None, None, None, None, ab2_full.index)
cute.copy(
tiled_copy_s2t_sfa, tCsSFA2_compact_s2t[s2t_coord2], tCtSFA2_compact_s2t
)
cute.copy(
tiled_copy_s2t_sfb, tCsSFB2_compact_s2t[s2t_coord2], tCtSFB2_compact_s2t
)
# MMA2: ACC2 += A2 @ B2
num_kblocks = cute.size(tCrA2, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, ab2_full.index)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(tcgen05.Field.SFA, tCtSFA2[sf_kblock_coord].iterator)
tiled_mma.set(tcgen05.Field.SFB, tCtSFB2[sf_kblock_coord].iterator)
cute.gemm(
tiled_mma,
tCtAcc2,
tCrA2[kblock_coord],
tCrB2[kblock_coord],
tCtAcc2,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab2_full.release()
acc2_empty.commit()
#
# EPILOGUE: All 4 warps participate
#
op = tcgen05.Ld32x32bOp(tcgen05.Repetition.x128, tcgen05.Pack.NONE)
copy_atom_t2r = cute.make_copy_atom(op, cutlass.Float32)
tiled_copy_t2r = tcgen05.make_tmem_copy(copy_atom_t2r, tCtAcc1)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tTR_tAcc1 = thr_copy_t2r.partition_S(tCtAcc1)
tTR_tAcc2 = thr_copy_t2r.partition_S(tCtAcc2)
tTR_gC = thr_copy_t2r.partition_D(tCgC)
tTR_rAcc1 = cute.make_rmem_tensor(
tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32
)
tTR_rAcc2 = cute.make_rmem_tensor(
tTR_gC[None, None, None, None, 0, 0, 0].shape, cutlass.Float32
)
tTR_rC = cute.make_rmem_tensor(
tTR_gC[None, None, None, None, 0, 0, 0].shape, c_dtype
)
simt_atom = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), c_dtype)
tTR_gC = tTR_gC[(None, None, None, None, *mma_tile_coord_mnl)]
acc1_full = acc1_consumer.wait_and_advance()
cute.copy(tiled_copy_t2r, tTR_tAcc1, tTR_rAcc1)
acc_vec1 = epilogue_op(tTR_rAcc1.load())
acc2_full = acc2_consumer.wait_and_advance()
cute.copy(tiled_copy_t2r, tTR_tAcc2, tTR_rAcc2)
acc_vec2 = tTR_rAcc2.load()
acc_vec = acc_vec1 * acc_vec2
tTR_rC.store(acc_vec.to(c_dtype))
# Store C to global memory
cute.copy(simt_atom, tTR_rC, tTR_gC)
acc1_full.release()
acc2_full.release()
# Deallocate TMEM
cute.arch.barrier()
tmem.free(acc_tmem_ptr)
return
@cute.jit
def my_kernel(
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: tuple,
epilogue_op: cutlass.Constexpr = lambda x: x
* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
m, n, k, l = problem_size
# A tensor (shared by both GEMMs, but we create two TMA descriptors)
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_tensor1 = cute.make_tensor(
b1_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
b_tensor2 = cute.make_tensor(
b2_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))
)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor1.shape, sf_vec_size)
sfb_tensor1 = cute.make_tensor(sfb1_ptr, sfb_layout)
sfb_tensor2 = cute.make_tensor(sfb2_ptr, sfb_layout)
mma_op = tcgen05.MmaMXF4NVF4Op(
sf_dtype,
(mma_tiler_mnk[0], mma_tiler_mnk[1], mma_inst_shape_k),
tcgen05.CtaGroup.ONE,
tcgen05.OperandSource.SMEM,
)
tiled_mma = cute.make_tiled_mma(mma_op)
cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout((1, 1, 1)), (tiled_mma.thr_id.shape,)
)
# Layouts (same for both GEMMs)
a_smem_layout_staged = sm100_utils.make_smem_layout_a(
tiled_mma, mma_tiler_mnk, ab_dtype, num_ab1_stage
)
b_smem_layout_staged = sm100_utils.make_smem_layout_b(
tiled_mma, mma_tiler_mnk, ab_dtype, num_ab1_stage
)
sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab1_stage
)
sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
tiled_mma, mma_tiler_mnk, sf_vec_size, num_ab1_stage
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
# TMA atoms for GEMM1 (A1, B1, SFA1, SFB1)
a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, None, 0))
b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, None, 0))
sfa_smem_layout = cute.slice_(sfa_smem_layout_staged, (None, None, None, 0))
sfb_smem_layout = cute.slice_(sfb_smem_layout_staged, (None, None, None, 0))
tma_atom_a1, tma_tensor_a1 = cute.nvgpu.make_tiled_tma_atom_A(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
a_tensor,
a_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
)
tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
b_tensor1,
b_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
)
tma_atom_sfa1, tma_tensor_sfa1 = cute.nvgpu.make_tiled_tma_atom_A(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
sfa_tensor,
sfa_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
sfb_tensor1,
sfb_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
# TMA atoms for GEMM2 (A2, B2, SFA2, SFB2) - same source tensors, different SMEM targets
tma_atom_a2, tma_tensor_a2 = cute.nvgpu.make_tiled_tma_atom_A(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
a_tensor,
a_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
)
tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
b_tensor2,
b_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
)
tma_atom_sfa2, tma_tensor_sfa2 = cute.nvgpu.make_tiled_tma_atom_A(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
sfa_tensor,
sfa_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE),
sfb_tensor2,
sfb_smem_layout,
mma_tiler_mnk,
tiled_mma,
cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
# Byte counts for each pipeline
a_copy_size = cute.size_in_bytes(ab_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(ab_dtype, b_smem_layout)
sfa_copy_size = cute.size_in_bytes(sf_dtype, sfa_smem_layout)
sfb_copy_size = cute.size_in_bytes(sf_dtype, sfb_smem_layout)
num_tma_bytes_gemm1 = (
a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
) * atom_thr_size
num_tma_bytes_gemm2 = (
a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
) * atom_thr_size
grid = (
cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),
cute.ceil_div(c_tensor.shape[1], mma_tiler_mnk[1]),
c_tensor.shape[2],
)
kernel(
tiled_mma,
# GEMM1
tma_atom_a1,
tma_tensor_a1,
tma_atom_b1,
tma_tensor_b1,
tma_atom_sfa1,
tma_tensor_sfa1,
tma_atom_sfb1,
tma_tensor_sfb1,
# GEMM2
tma_atom_a2,
tma_tensor_a2,
tma_atom_b2,
tma_tensor_b2,
tma_atom_sfa2,
tma_tensor_sfa2,
tma_atom_sfb2,
tma_tensor_sfb2,
# Output
c_tensor,
# Layouts
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
# Byte counts
num_tma_bytes_gemm1,
num_tma_bytes_gemm2,
epilogue_op,
).launch(
grid=grid,
block=[threads_per_cta, 1, 1],
cluster=(1, 1, 1),
)
return
# Global cache for compiled kernel
_compiled_kernel_cache = None
def compile_kernel():
global _compiled_kernel_cache
if _compiled_kernel_cache is not None:
return _compiled_kernel_cache
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)
_compiled_kernel_cache = cute.compile(
my_kernel,
a_ptr,
b1_ptr,
b2_ptr,
sfa_ptr,
sfb1_ptr,
sfb2_ptr,
c_ptr,
(0, 0, 0, 0),
)
return _compiled_kernel_cache
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
compiled_func = compile_kernel()
_, k, _ = a.shape
m, n, l = c.shape
k = k * 2
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
)
compiled_func(
a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, (m, n, k, l)
)
return c
scrolls · 907 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON