submission 488633
yolojewjitsu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 890 lines, June 9 Researcher Reciprocity License v1.0.
mission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-488633?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:5ab8798f35d3d486c28c7042584ddcc4709daf1432ddfa7db2e8baec74f570ae
license declaredunknown
license concludedunknown
authorsyolojewjitsu
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(epilog_warp_ids))persistent-kernel
tile_sched_params: utils.PersistentTileSchedulerParams,shared-memory
smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")tcgen05
cta_group = tcgen05.CtaGroup.ONEwarp-specialization
num_tma_producer = num_mcast_ctas_a + num_mcast_ctas_b - 1Kernel source
mission.py890 lines
import functools
from typing import List, Tuple, Union
from inspect import isclass
import torch
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import make_ptr
from task import input_t, output_t
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
acc_dtype = cutlass.Float32
sf_vec_size = 16
mma_tiler_mn = (128, 128)
cluster_shape_mn = (1, 2)
mma_inst_shape_k = 64
mma_inst_tile_k = 4
mma_tiler_mnk = (*mma_tiler_mn, mma_inst_shape_k * mma_inst_tile_k)
bytes_per_tensormap = 128
num_tensormaps = 5
epilog_warp_ids = (0, 1, 2, 3)
mma_warp_id = 4
tma_warp_id = 5
threads_per_cta = 32 * 6
epilog_sync_barrier = pipeline.NamedBarrier(barrier_id=1, num_threads=32 * len(epilog_warp_ids))
tmem_alloc_barrier = pipeline.NamedBarrier(barrier_id=2, num_threads=32 * (1 + len(epilog_warp_ids)))
tensormap_ab_init_barrier = pipeline.NamedBarrier(barrier_id=3, num_threads=64)
SM100_TMEM_CAPACITY_COLUMNS = 512
smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
def compute_stages(c_layout):
"""Compute num_acc_stage, num_ab_stage, num_c_stage."""
cta_group = tcgen05.CtaGroup.ONE
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
ab_dtype, "M", "N", sf_dtype, sf_vec_size, cta_group, mma_tiler_mn,
)
num_acc_stage = 2
epi_tile = sm100_utils.compute_epilogue_tile_shape(
(mma_tiler_mnk[0], mma_tiler_mnk[1], mma_tiler_mnk[2]),
False,
c_layout,
c_dtype,
)
num_c_stage = 2
a_smem_1 = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, ab_dtype, 1)
b_smem_1 = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, ab_dtype, 1)
sfa_smem_1 = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
sfb_smem_1 = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, 1)
c_smem_1 = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, 1)
ab_bytes = (
cute.size_in_bytes(ab_dtype, a_smem_1)
+ cute.size_in_bytes(ab_dtype, b_smem_1)
+ cute.size_in_bytes(sf_dtype, sfa_smem_1)
+ cute.size_in_bytes(sf_dtype, sfb_smem_1)
)
c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_1)
reserved = 1024
occupancy = 1
c_bytes = c_bytes_per_stage * num_c_stage
num_ab_stage = (smem_capacity // occupancy - (reserved + c_bytes)) // ab_bytes
num_c_stage += (
smem_capacity - occupancy * ab_bytes * num_ab_stage
- occupancy * (reserved + c_bytes)
) // (occupancy * c_bytes_per_stage)
return num_acc_stage, num_ab_stage, num_c_stage, epi_tile
@cute.jit
def make_real_tensor_ab(ptr_i64, dtype, outer_dim, k):
"""Construct global tensor for A (m,k,1) or B (n,k,1) from pointer and dimensions."""
tensor_ptr = cute.make_ptr(dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
c1 = cutlass.Int32(1)
c0 = cutlass.Int32(0)
return cute.make_tensor(
tensor_ptr,
cute.make_layout(
(outer_dim, k, c1),
stride=(k, cutlass.Int32(1), c0),
),
)
@cute.jit
def make_real_tensor_c(ptr_i64, m, n):
"""Construct global tensor for C from pointer and dimensions."""
tensor_ptr = cute.make_ptr(c_dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
c1 = cutlass.Int32(1)
c0 = cutlass.Int32(0)
return cute.make_tensor(
tensor_ptr,
cute.make_layout(
(m, n, c1),
stride=(n, cutlass.Int32(1), c0),
),
)
@cute.jit
def make_real_tensor_sf(ptr_i64, outer_dim, k):
"""Construct global tensor for SFA or SFB from pointer and dimensions."""
tensor_ptr = cute.make_ptr(sf_dtype, ptr_i64, cute.AddressSpace.gmem, assumed_align=16)
c1 = cutlass.Int32(1)
sf_layout = blockscaled_utils.tile_atom_to_shape_SF(
(outer_dim, k, c1), sf_vec_size
)
return cute.make_tensor(tensor_ptr, sf_layout)
@cute.kernel
def kernel(
tiled_mma: 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,
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,
group_count: cutlass.Constexpr,
problem_sizes_mnkl: cute.Tensor,
ptrs_abc: cute.Tensor,
ptrs_sfasfb: cute.Tensor,
tensormaps: cute.Tensor,
num_tma_load_bytes: cutlass.Constexpr[int],
num_ab_stage: cutlass.Constexpr[int],
num_acc_stage: cutlass.Constexpr[int],
num_c_stage: cutlass.Constexpr[int],
cluster_tile_shape_mnk: cutlass.Constexpr,
c_layout: cutlass.Constexpr,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
if warp_idx == 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)
bidx, bidy, bidz = cute.arch.block_idx()
tidx, _, _ = cute.arch.thread_idx()
size_tensormap_in_i64 = num_tensormaps * bytes_per_tensormap // 8
@cute.struct
class SharedStorage:
tensormap_buffer: cute.struct.MemRange[cutlass.Int64, size_tensormap_in_i64]
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[c_dtype, cute.cosize(c_smem_layout_staged.outer)],
1024,
]
sA: cute.struct.Align[
cute.struct.MemRange[ab_dtype, cute.cosize(a_smem_layout_staged.outer)],
1024,
]
sB: cute.struct.Align[
cute.struct.MemRange[ab_dtype, cute.cosize(b_smem_layout_staged.outer)],
1024,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[sf_dtype, cute.cosize(sfa_smem_layout_staged)],
1024,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[sf_dtype, cute.cosize(sfb_smem_layout_staged)],
1024,
]
smem = utils.SmemAllocator()
storage = smem.allocate(SharedStorage)
tensormap_smem_ptr = storage.tensormap_buffer.data_ptr()
tm_step = bytes_per_tensormap // 8
tensormap_a_smem_ptr = tensormap_smem_ptr
tensormap_b_smem_ptr = tensormap_a_smem_ptr + tm_step
tensormap_sfa_smem_ptr = tensormap_b_smem_ptr + tm_step
tensormap_sfb_smem_ptr = tensormap_sfa_smem_ptr + tm_step
tensormap_c_smem_ptr = tensormap_sfb_smem_ptr + tm_step
tmem_holding_buf = storage.tmem_holding_buf
cluster_layout_vmnk = cute.tiled_divide(cute.make_layout((*cluster_shape_mn, 1)), (tiled_mma.thr_id.shape,))
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)
num_mcast_ctas_a = cute.size(cluster_layout_vmnk.shape[2])
num_mcast_ctas_b = cute.size(cluster_layout_vmnk.shape[1])
num_tma_producer = num_mcast_ctas_a + num_mcast_ctas_b - 1
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=num_ab_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, num_tma_producer),
tx_count=num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
)
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=num_acc_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread),
consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, len(epilog_warp_ids)),
cta_layout_vmnk=cluster_layout_vmnk,
)
pipeline_init_arrive(cluster_shape_mn=cluster_shape_mn, is_relaxed=True)
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)
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_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1)
gA_mkl = cute.local_tile(mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None))
gB_nkl = cute.local_tile(mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None))
gSFA_mkl = cute.local_tile(mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None))
gSFB_nkl = cute.local_tile(mSFB_nkl, 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))
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
thr_mma = tiled_mma.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.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
a_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape)
b_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, None, 0, 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))
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))
tAsSFA, tAgSFA = cpasync.tma_partition(tma_atom_sfa, block_in_cluster_coord_vmnk[2], a_cta_layout, cute.group_modes(sSFA, 0, 3), cute.group_modes(tCgSFA, 0, 3))
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
tBsSFB, tBgSFB = cpasync.tma_partition(tma_atom_sfb, block_in_cluster_coord_vmnk[1], b_cta_layout, cute.group_modes(sSFB, 0, 3), cute.group_modes(tCgSFB, 0, 3))
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
tCrA = tiled_mma.make_fragment_A(sA)
tCrB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(mma_tiler_mn)
tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, num_acc_stage))
pipeline_init_wait(cluster_shape_mn=cluster_shape_mn)
grid_dim = cute.arch.grid_dim()
tensormap_workspace_idx = bidz * grid_dim[1] * grid_dim[0] + bidy * grid_dim[0] + bidx
tensormap_manager = utils.TensorMapManager(utils.TensorMapUpdateMode.SMEM, bytes_per_tensormap)
tensormap_a_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 0, None)].iterator)
tensormap_b_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 1, None)].iterator)
tensormap_sfa_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 2, None)].iterator)
tensormap_sfb_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 3, None)].iterator)
tensormap_c_gmem_ptr = tensormap_manager.get_tensormap_ptr(tensormaps[(tensormap_workspace_idx, 4, None)].iterator)
if warp_idx == tma_warp_id:
tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
)
tensormap_init_done = cutlass.Boolean(False)
last_group_idx = cutlass.Int32(-1)
work_tile = tile_sched.initial_work_tile_info()
ab_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, num_ab_stage)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
grouped_gemm_info = group_gemm_ts_helper.delinearize_z(cur_tile_coord, problem_sizes_mnkl)
cur_k_tile_cnt = grouped_gemm_info.cta_tile_count_k
cur_group_idx = grouped_gemm_info.group_idx
is_group_changed = cur_group_idx != last_group_idx
if is_group_changed:
m_cur = grouped_gemm_info.problem_shape_m
n_cur = grouped_gemm_info.problem_shape_n
k_cur = grouped_gemm_info.problem_shape_k
real_a = make_real_tensor_ab(ptrs_abc[(cur_group_idx, 0)], ab_dtype, m_cur, k_cur)
real_b = make_real_tensor_ab(ptrs_abc[(cur_group_idx, 1)], ab_dtype, n_cur, k_cur)
real_sfa = make_real_tensor_sf(ptrs_sfasfb[(cur_group_idx, 0)], m_cur, k_cur)
real_sfb = make_real_tensor_sf(ptrs_sfasfb[(cur_group_idx, 1)], n_cur, k_cur)
if tensormap_init_done == False:
tensormap_ab_init_barrier.arrive_and_wait()
tensormap_init_done = True
tensormap_manager.update_tensormap(
(real_a, real_b, real_sfa, real_sfb),
(tma_atom_a, tma_atom_b, tma_atom_sfa, tma_atom_sfb),
(tensormap_a_gmem_ptr, tensormap_b_gmem_ptr, tensormap_sfa_gmem_ptr, tensormap_sfb_gmem_ptr),
tma_warp_id,
(tensormap_a_smem_ptr, tensormap_b_smem_ptr, tensormap_sfa_smem_ptr, tensormap_sfb_smem_ptr),
)
mma_tile_coord_mnl = (
grouped_gemm_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
grouped_gemm_info.cta_tile_idx_n,
0,
)
tAgA_s = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB_s = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA_s = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgSFB_s = tBgSFB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
ab_producer_state.reset_count()
peek_status = cutlass.Boolean(1)
if ab_producer_state.count < cur_k_tile_cnt:
peek_status = ab_pipeline.producer_try_acquire(ab_producer_state)
if is_group_changed:
tensormap_manager.fence_tensormap_update(tensormap_a_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_b_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_sfa_gmem_ptr)
tensormap_manager.fence_tensormap_update(tensormap_sfb_gmem_ptr)
for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_status)
cute.copy(tma_atom_a, tAgA_s[(None, ab_producer_state.count)], tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_a_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_b, tBgB_s[(None, ab_producer_state.count)], tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_b_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_sfa, tAgSFA_s[(None, ab_producer_state.count)], tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfa_gmem_ptr, cute.AddressSpace.generic))
cute.copy(tma_atom_sfb, tBgSFB_s[(None, ab_producer_state.count)], tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_sfb_gmem_ptr, cute.AddressSpace.generic))
ab_producer_state.advance()
peek_status = cutlass.Boolean(1)
if ab_producer_state.count < cur_k_tile_cnt:
peek_status = ab_pipeline.producer_try_acquire(ab_producer_state)
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
last_group_idx = cur_group_idx
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == mma_warp_id:
tensormap_manager.init_tensormap_from_atom(tma_atom_a, tensormap_a_smem_ptr, mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_b, tensormap_b_smem_ptr, mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_sfa, tensormap_sfa_smem_ptr, mma_warp_id)
tensormap_manager.init_tensormap_from_atom(tma_atom_sfb, tensormap_sfb_smem_ptr, mma_warp_id)
tensormap_ab_init_barrier.arrive_and_wait()
tmem_alloc_barrier.arrive_and_wait()
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
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=sf_dtype)
tCtSFA_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)))
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=sf_dtype)
tCtSFB_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)))
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
copy_atom_s2t = cute.make_copy_atom(tcgen05.Cp4x32x128bOp(tcgen05.CtaGroup.ONE), sf_dtype)
tCsSFA_compact = cute.filter_zeros(sSFA)
tCtSFA_compact = cute.filter_zeros(tCtSFA)
tiled_copy_s2t_sfa = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFA_compact)
thr_s2t_sfa = tiled_copy_s2t_sfa.get_slice(0)
tCsSFA_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t_sfa, thr_s2t_sfa.partition_S(tCsSFA_compact))
tCtSFA_s2t = thr_s2t_sfa.partition_D(tCtSFA_compact)
tCsSFB_compact = cute.filter_zeros(sSFB)
tCtSFB_compact = cute.filter_zeros(tCtSFB)
tiled_copy_s2t_sfb = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSFB_compact)
thr_s2t_sfb = tiled_copy_s2t_sfb.get_slice(0)
tCsSFB_s2t = tcgen05.get_s2t_smem_desc_tensor(tiled_copy_s2t_sfb, thr_s2t_sfb.partition_S(tCsSFB_compact))
tCtSFB_s2t = thr_s2t_sfb.partition_D(tCtSFB_compact)
tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
)
work_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, num_ab_stage)
acc_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, num_acc_stage)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
cur_k_tile_cnt, cur_group_idx = group_gemm_ts_helper.search_cluster_tile_count_k(cur_tile_coord, problem_sizes_mnkl)
tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
ab_consumer_state.reset_count()
peek_status = cutlass.Boolean(1)
if ab_consumer_state.count < cur_k_tile_cnt:
peek_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
acc_pipeline.producer_acquire(acc_producer_state)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(cur_k_tile_cnt):
ab_pipeline.consumer_wait(ab_consumer_state, peek_status)
s2t_coord = (None, None, None, None, ab_consumer_state.index)
cute.copy(tiled_copy_s2t_sfa, tCsSFA_s2t[s2t_coord], tCtSFA_s2t)
cute.copy(tiled_copy_s2t_sfb, tCsSFB_s2t[s2t_coord], tCtSFB_s2t)
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, ab_consumer_state.index)
sf_coord = (None, None, kblock_idx)
tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_coord].iterator)
tiled_mma.set(tcgen05.Field.SFB, tCtSFB[sf_coord].iterator)
cute.gemm(tiled_mma, tCtAcc, tCrA[kblock_coord], tCrB[kblock_coord], tCtAcc)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_status = cutlass.Boolean(1)
if ab_consumer_state.count < cur_k_tile_cnt:
peek_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
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)
if warp_idx < mma_warp_id:
tensormap_manager.init_tensormap_from_atom(tma_atom_c, tensormap_c_smem_ptr, epilog_warp_ids[0])
if warp_idx == epilog_warp_ids[0]:
cute.arch.alloc_tmem(SM100_TMEM_CAPACITY_COLUMNS, tmem_holding_buf, is_two_cta=False)
tmem_alloc_barrier.arrive_and_wait()
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(acc_dtype, alignment=16, ptr_to_buffer_holding_addr=tmem_holding_buf)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epi_tidx = tidx
copy_atom_t2r = sm100_utils.get_tmem_load_op(
(mma_tiler_mnk[0], mma_tiler_mnk[1], mma_tiler_mnk[2]),
c_layout, c_dtype, acc_dtype, epi_tile, False,
)
tAcc_epi = cute.flat_divide(tCtAcc_base[((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(epi_tidx)
tTR_tAcc_base = thr_copy_t2r.partition_S(tAcc_epi)
gC_mnl_epi = cute.flat_divide(tCgC[((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, acc_dtype)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, c_dtype)
copy_atom_r2s = sm100_utils.get_smem_store_op(c_layout, c_dtype, 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(epi_tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
sC_for_tma = cute.group_modes(sC, 0, 2)
gC_for_tma = cute.group_modes(gC_mnl_epi, 0, 2)
bSG_sC, bSG_gC_partitioned = cpasync.tma_partition(tma_atom_c, 0, cute.make_layout(1), sC_for_tma, gC_for_tma)
tile_sched = utils.StaticPersistentTileScheduler.create(tile_sched_params, cute.arch.block_idx(), grid_dim)
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count, tile_sched_params, cluster_tile_shape_mnk, utils.create_initial_search_state(),
)
work_tile = tile_sched.initial_work_tile_info()
acc_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, num_acc_stage)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=num_c_stage,
producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 32 * len(epilog_warp_ids)),
)
last_group_idx = cutlass.Int32(-1)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
grouped_gemm_info = group_gemm_ts_helper.delinearize_z(cur_tile_coord, problem_sizes_mnkl)
cur_group_idx = grouped_gemm_info.group_idx
is_group_changed = cur_group_idx != last_group_idx
if is_group_changed:
m_cur = grouped_gemm_info.problem_shape_m
n_cur = grouped_gemm_info.problem_shape_n
k_cur = grouped_gemm_info.problem_shape_k
real_c = make_real_tensor_c(ptrs_abc[(cur_group_idx, 2)], m_cur, n_cur)
tensormap_manager.update_tensormap(
((real_c),), ((tma_atom_c),), ((tensormap_c_gmem_ptr),),
epilog_warp_ids[0], (tensormap_c_smem_ptr,),
)
mma_tile_coord_mnl = (
grouped_gemm_info.cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
grouped_gemm_info.cta_tile_idx_n,
0,
)
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))
if is_group_changed:
if warp_idx == epilog_warp_ids[0]:
tensormap_manager.fence_tensormap_update(tensormap_c_gmem_ptr)
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in range(subtile_cnt):
cute.copy(tiled_copy_t2r, tTR_tAcc[(None, None, None, subtile_idx)], tTR_rAcc)
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
tRS_rC.store(acc_vec.to(c_dtype))
c_buffer = (num_prev_subtiles + subtile_idx) % num_c_stage
cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
cute.arch.fence_proxy("async.shared", space="cta")
epilog_sync_barrier.arrive_and_wait()
if warp_idx == epilog_warp_ids[0]:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, subtile_idx)],
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(tensormap_c_gmem_ptr, cute.AddressSpace.generic),
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
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()
last_group_idx = cur_group_idx
if warp_idx == epilog_warp_ids[0]:
cute.arch.relinquish_tmem_alloc_permit(is_two_cta=False)
epilog_sync_barrier.arrive_and_wait()
if warp_idx == epilog_warp_ids[0]:
cute.arch.dealloc_tmem(acc_tmem_ptr, SM100_TMEM_CAPACITY_COLUMNS, is_two_cta=False)
c_pipeline.producer_tail()
@cute.jit
def launch_kernel(
ptr_problem_sizes: cute.Pointer,
ptr_abc: cute.Pointer,
ptr_sfasfb: cute.Pointer,
ptr_tensormap: cute.Pointer,
total_num_clusters: cutlass.Constexpr[int],
num_tensormap_entries: cutlass.Constexpr[int],
problem_sizes: List[Tuple[int, int, int, int]],
max_active_clusters: cutlass.Constexpr[int],
):
min_shape = (cutlass.Int32(128), cutlass.Int32(128), cutlass.Int32(256), cutlass.Int32(1))
null_ab = cute.make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
null_sf = cute.make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
null_c = cute.make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
initial_a = cute.make_tensor(null_ab, cute.make_layout(
(min_shape[0], cute.assume(min_shape[2], 32), min_shape[3]),
stride=(cute.assume(min_shape[2], 32), 1, cute.assume(min_shape[0] * min_shape[2], 32)),
))
initial_b = cute.make_tensor(null_ab, cute.make_layout(
(min_shape[1], cute.assume(min_shape[2], 32), min_shape[3]),
stride=(cute.assume(min_shape[2], 32), 1, cute.assume(min_shape[1] * min_shape[2], 32)),
))
initial_c = cute.make_tensor(null_c, cute.make_layout(
(cute.assume(min_shape[0], 32), min_shape[1], min_shape[3]),
stride=(min_shape[1], 1, cute.assume(min_shape[0] * min_shape[1], 32)),
))
c_layout = utils.LayoutEnum.from_tensor(initial_c)
NUM_ACC_STAGE, NUM_AB_STAGE, NUM_C_STAGE, _ = compute_stages(c_layout)
num_groups = len(problem_sizes)
tensor_problem_sizes = cute.make_tensor(ptr_problem_sizes, cute.make_layout((num_groups, 4), stride=(4, 1)))
tensor_abc = cute.make_tensor(ptr_abc, cute.make_layout((num_groups, 3), stride=(3, 1)))
tensor_sfasfb = cute.make_tensor(ptr_sfasfb, cute.make_layout((num_groups, 2), stride=(2, 1)))
tensor_tensormap = cute.make_tensor(ptr_tensormap, cute.make_layout((num_tensormap_entries, num_tensormaps, bytes_per_tensormap // 8), stride=(num_tensormaps * bytes_per_tensormap // 8, bytes_per_tensormap // 8, 1)))
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_a.shape, sf_vec_size)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(initial_b.shape, sf_vec_size)
initial_sfa = cute.make_tensor(null_sf, sfa_layout)
initial_sfb = cute.make_tensor(null_sf, sfb_layout)
cta_group = tcgen05.CtaGroup.ONE
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
ab_dtype, "M", "N", sf_dtype, sf_vec_size, cta_group, mma_tiler_mn,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
cluster_layout_vmnk = cute.tiled_divide(cute.make_layout((*cluster_shape_mn, 1)), (tiled_mma.thr_id.shape,))
is_a_mcast = cluster_shape_mn[1] > 1
is_b_mcast = cluster_shape_mn[0] > 1
tma_op_a = cpasync.CopyBulkTensorTileG2SMulticastOp(cta_group) if is_a_mcast else cpasync.CopyBulkTensorTileG2SOp(cta_group)
tma_op_b = cpasync.CopyBulkTensorTileG2SMulticastOp(cta_group) if is_b_mcast else cpasync.CopyBulkTensorTileG2SOp(cta_group)
a_smem_staged = sm100_utils.make_smem_layout_a(tiled_mma, mma_tiler_mnk, ab_dtype, NUM_AB_STAGE)
b_smem_staged = sm100_utils.make_smem_layout_b(tiled_mma, mma_tiler_mnk, ab_dtype, NUM_AB_STAGE)
sfa_smem_staged = blockscaled_utils.make_smem_layout_sfa(tiled_mma, mma_tiler_mnk, sf_vec_size, NUM_AB_STAGE)
sfb_smem_staged = blockscaled_utils.make_smem_layout_sfb(tiled_mma, mma_tiler_mnk, sf_vec_size, NUM_AB_STAGE)
epi_tile = sm100_utils.compute_epilogue_tile_shape(
mma_tiler_mnk, False, c_layout, c_dtype,
)
c_smem_staged = sm100_utils.make_smem_layout_epi(c_dtype, c_layout, epi_tile, NUM_C_STAGE)
a_smem = cute.slice_(a_smem_staged, (None, None, None, 0))
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
tma_op_a, initial_a, a_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
)
b_smem = cute.slice_(b_smem_staged, (None, None, None, 0))
tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
tma_op_b, initial_b, b_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
)
sfa_smem = cute.slice_(sfa_smem_staged, (None, None, None, 0))
tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
tma_op_a, initial_sfa, sfa_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
sfb_smem = cute.slice_(sfb_smem_staged, (None, None, None, 0))
tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
tma_op_b, initial_sfb, sfb_smem, mma_tiler_mnk, tiled_mma, cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
epi_smem = cute.slice_(c_smem_staged, (None, None, 0))
tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(), initial_c, epi_smem, epi_tile,
)
num_tma_load_bytes = (
cute.size_in_bytes(ab_dtype, a_smem)
+ cute.size_in_bytes(ab_dtype, b_smem)
+ cute.size_in_bytes(sf_dtype, sfa_smem)
+ cute.size_in_bytes(sf_dtype, sfb_smem)
) * atom_thr_size
size_tensormap_in_i64 = num_tensormaps * bytes_per_tensormap // 8
@cute.struct
class LauncherSharedStorage:
tensormap_buffer: cute.struct.MemRange[cutlass.Int64, size_tensormap_in_i64]
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_AB_STAGE]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_AB_STAGE]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_ACC_STAGE]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, NUM_ACC_STAGE]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[c_dtype, cute.cosize(c_smem_staged.outer)],
1024,
]
sA: cute.struct.Align[
cute.struct.MemRange[ab_dtype, cute.cosize(a_smem_staged.outer)],
1024,
]
sB: cute.struct.Align[
cute.struct.MemRange[ab_dtype, cute.cosize(b_smem_staged.outer)],
1024,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[sf_dtype, cute.cosize(sfa_smem_staged)],
1024,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[sf_dtype, cute.cosize(sfb_smem_staged)],
1024,
]
cta_tile_shape_mnk = (mma_tiler_mnk[0] // cute.size(tiled_mma.thr_id.shape), mma_tiler_mnk[1], mma_tiler_mnk[2])
cluster_tile_shape_mnk = tuple(x * y for x, y in zip(cta_tile_shape_mnk, (*cluster_shape_mn, 1)))
problem_shape_ntile_mnl = (cluster_shape_mn[0], cluster_shape_mn[1], cutlass.Int32(total_num_clusters))
tile_sched_params = utils.PersistentTileSchedulerParams(problem_shape_ntile_mnl, (*cluster_shape_mn, 1))
grid = utils.StaticPersistentTileScheduler.get_grid_shape(tile_sched_params, max_active_clusters)
kernel(
tiled_mma, 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,
a_smem_staged, b_smem_staged, sfa_smem_staged, sfb_smem_staged,
c_smem_staged, epi_tile,
tile_sched_params, len(problem_sizes), tensor_problem_sizes,
tensor_abc, tensor_sfasfb, tensor_tensormap,
num_tma_load_bytes, NUM_AB_STAGE, NUM_ACC_STAGE, NUM_C_STAGE,
cluster_tile_shape_mnk, c_layout,
).launch(
grid=grid,
block=[threads_per_cta, 1, 1],
cluster=(*cluster_shape_mn, 1),
smem=LauncherSharedStorage.size_in_bytes(),
min_blocks_per_mp=1,
)
_kernel_cache = {}
def _fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
abc_cpu, sfasfb_cpu, psizes_cpu, num_groups):
"""Fill pinned CPU buffers with current data pointers and problem sizes."""
for i in range(num_groups):
a, b, c = abc_tensors[i]
sfa, sfb = sfasfb_reordered_tensors[i]
abc_cpu[i, 0] = a.data_ptr()
abc_cpu[i, 1] = b.data_ptr()
abc_cpu[i, 2] = c.data_ptr()
sfasfb_cpu[i, 0] = sfa.data_ptr()
sfasfb_cpu[i, 1] = sfb.data_ptr()
psizes_cpu[i, 0] = problem_sizes[i][0]
psizes_cpu[i, 1] = problem_sizes[i][1]
psizes_cpu[i, 2] = problem_sizes[i][2]
psizes_cpu[i, 3] = problem_sizes[i][3]
def custom_kernel(data: input_t) -> output_t:
abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
num_groups = len(problem_sizes)
cluster_tile_m = mma_tiler_mnk[0] * cluster_shape_mn[0]
cluster_tile_n = mma_tiler_mnk[1] * cluster_shape_mn[1]
total_num_clusters = 0
for m, n, k, l in problem_sizes:
total_num_clusters += -(-m // cluster_tile_m) * -(-n // cluster_tile_n)
cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]
num_tensormap_entries = total_num_clusters * cluster_size
cache_key = (num_groups, total_num_clusters)
if cache_key not in _kernel_cache:
# ---- Compile kernel ----
null_ptr = make_ptr(cutlass.Int32, 0, cute.AddressSpace.gmem, assumed_align=16)
null_i64 = make_ptr(cutlass.Int64, 0, cute.AddressSpace.gmem, assumed_align=16)
hardware_info = cutlass.utils.HardwareInfo()
max_active_clusters = hardware_info.get_max_active_clusters(cluster_size)
dummy_sizes = [(128, 128, 256, 1)] * num_groups
compiled = cute.compile(
launch_kernel,
null_ptr, null_i64, null_i64, null_i64,
total_num_clusters,
num_tensormap_entries,
dummy_sizes,
max_active_clusters,
options="--opt-level 2",
)
# ---- Pre-allocate GPU + pinned CPU buffers ----
t_problem_sizes = torch.empty((num_groups, 4), dtype=torch.int32, device="cuda")
t_tensormap = torch.empty(
(num_tensormap_entries, num_tensormaps, bytes_per_tensormap // 8),
dtype=torch.int64, device="cuda",
)
t_abc_ptrs = torch.empty((num_groups, 3), dtype=torch.int64, device="cuda")
t_sfasfb_ptrs = torch.empty((num_groups, 2), dtype=torch.int64, device="cuda")
abc_cpu = torch.empty((num_groups, 3), dtype=torch.int64, pin_memory=True)
sfasfb_cpu = torch.empty((num_groups, 2), dtype=torch.int64, pin_memory=True)
psizes_cpu = torch.empty((num_groups, 4), dtype=torch.int32, pin_memory=True)
p_sizes = make_ptr(cutlass.Int32, t_problem_sizes.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
p_abc = make_ptr(cutlass.Int64, t_abc_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
p_sfasfb = make_ptr(cutlass.Int64, t_sfasfb_ptrs.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
p_tmap = make_ptr(cutlass.Int64, t_tensormap.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
# ---- Warmup call (produces correct results for first invocation) ----
_fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
abc_cpu, sfasfb_cpu, psizes_cpu, num_groups)
t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
torch.cuda.synchronize()
# ---- Try CUDA Graph capture (H2D copies + kernel in one graph) ----
cuda_graph = None
try:
cuda_graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(cuda_graph):
t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
except Exception:
cuda_graph = None
# Verify graph actually captured the kernel (not just H2D copies)
if cuda_graph is not None:
abc_tensors[0][2].fill_(0)
cuda_graph.replay()
torch.cuda.synchronize()
if abc_tensors[0][2].abs().max().item() < 1e-6:
cuda_graph = None # Graph didn't capture kernel launch
# Re-run kernel directly to restore correct C values
compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
torch.cuda.synchronize()
_kernel_cache[cache_key] = (
compiled, cuda_graph,
t_abc_ptrs, t_sfasfb_ptrs, t_problem_sizes,
abc_cpu, sfasfb_cpu, psizes_cpu,
p_sizes, p_abc, p_sfasfb, p_tmap,
t_tensormap,
)
# First call: warmup (or verification replay) produced correct results
return [abc_tensors[i][2] for i in range(num_groups)]
(compiled, cuda_graph,
t_abc_ptrs, t_sfasfb_ptrs, t_problem_sizes,
abc_cpu, sfasfb_cpu, psizes_cpu,
p_sizes, p_abc, p_sfasfb, p_tmap, _) = _kernel_cache[cache_key]
# Fill pinned CPU buffers (CPU-only, no CUDA calls)
_fill_pinned_buffers(abc_tensors, sfasfb_reordered_tensors, problem_sizes,
abc_cpu, sfasfb_cpu, psizes_cpu, num_groups)
if cuda_graph is not None:
# CUDA Graph replay: H2D copies + kernel launch in ~3-5μs
cuda_graph.replay()
else:
# Fallback: individual copies + dispatch
t_abc_ptrs.copy_(abc_cpu, non_blocking=True)
t_sfasfb_ptrs.copy_(sfasfb_cpu, non_blocking=True)
t_problem_sizes.copy_(psizes_cpu, non_blocking=True)
compiled(p_sizes, p_abc, p_sfasfb, p_tmap, problem_sizes)
return [abc_tensors[i][2] for i in range(num_groups)]
scrolls · 890 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