submission 188821
jason · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1396 lines, June 9 Researcher Reciprocity License v1.0.
submission8.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-188821?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:3333703681ee1b4b336755c98c83b5ebcfc0fb6c1985f49e1b013e77d5fbeb1a
license declaredunknown
license concludedunknown
authorsjason
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(persistent-kernel
class PersistentNvfp4Gemm: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
submission8.py1396 lines
from dataclasses import dataclass
import torch
import cutlass
import cutlass.cute as cute
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
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
from task import input_t, output_t
@dataclass(frozen=True)
class KernelConfig:
name: str
mma_tiler_mn: tuple
cluster_shape_mn: tuple
swizzle_size: int
raster_along_m: bool
occupancy: int
CONFIGS = {
"n7168_k16384": KernelConfig(
name="n7168_k16384",
mma_tiler_mn=(256, 128),
cluster_shape_mn=(2, 1),
swizzle_size=4,
raster_along_m=True,
occupancy=1,
),
"n4096_k7168": KernelConfig(
name="n4096_k7168",
mma_tiler_mn=(256, 128),
cluster_shape_mn=(2, 1),
swizzle_size=1,
raster_along_m=False,
occupancy=1,
),
"n7168_k2048": KernelConfig(
name="n7168_k2048",
mma_tiler_mn=(128, 192),
cluster_shape_mn=(1, 1),
swizzle_size=2,
raster_along_m=False,
occupancy=1,
),
"default": KernelConfig(
name="default",
mma_tiler_mn=(128, 128),
cluster_shape_mn=(1, 1),
swizzle_size=1,
raster_along_m=True,
occupancy=1,
),
}
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
sf_vec_size = 16
class PersistentNvfp4Gemm:
def __init__(self, cfg: KernelConfig):
self.sf_vec_size = sf_vec_size
self.mma_tiler_mn = cfg.mma_tiler_mn
self.cluster_shape_mn = cfg.cluster_shape_mn
self.swizzle_size = cfg.swizzle_size
self.raster_along_m = cfg.raster_along_m
self.occupancy = cfg.occupancy
self.acc_dtype = cutlass.Float32
self.use_2cta_instrs = self.mma_tiler_mn[0] == 256
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.epilog_warp_id = (0, 1, 2, 3)
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_cta = 32 * len(
(self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
)
self.epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=1, num_threads=32 * len(self.epilog_warp_id)
)
self.tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=2, num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id))
)
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
self.num_tmem_alloc_cols = 512
def _setup_attributes(self, a_tensor, b_tensor, c_tensor):
self.a_dtype = a_tensor.element_type
self.b_dtype = b_tensor.element_type
self.sf_dtype = sf_dtype
self.c_dtype = 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_tensor).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(c_tensor)
self.mma_inst_shape_mn = (self.mma_tiler_mn[0], self.mma_tiler_mn[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,
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
)
self.overlapping_accum = self.num_acc_stage == 1
sf_atom_mn = 32
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
if not self.overlapping_accum:
self.num_accumulator_tmem_cols = (
self.cta_tile_shape_mnk[1] * self.num_acc_stage
)
else:
self.num_accumulator_tmem_cols = (
self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols
)
self.iter_acc_early_release_in_epilogue = self.num_sf_tmem_cols // self.epi_tile_n
@cute.jit
def __call__(
self,
a_tensor: cute.Tensor,
b_tensor: cute.Tensor,
sfa_tensor: cute.Tensor,
sfb_tensor: cute.Tensor,
c_tensor: cute.Tensor,
max_active_clusters: cutlass.Constexpr,
):
self._setup_attributes(a_tensor, b_tensor, c_tensor)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
a_tensor.shape, self.sf_vec_size
)
sfa_tensor = cute.make_tensor(sfa_tensor.iterator, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
b_tensor.shape, self.sf_vec_size
)
sfb_tensor = cute.make_tensor(sfb_tensor.iterator, 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,
tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
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_op_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
b_op = sm100_utils.cluster_shape_to_tma_atom_B(
self.cluster_shape_mn, tiled_mma.thr_id
)
b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
tma_op_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0))
tma_op_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0))
tma_op_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x = tma_tensor_sfb.stride[0][1]
y = cute.ceil_div(tma_tensor_sfb.shape[0][1], 4)
new_shape = (
(tma_tensor_sfb.shape[0][0], ((2, 2), y)),
tma_tensor_sfb.shape[1],
tma_tensor_sfb.shape[2],
)
x_times_3 = 3 * x
new_stride = (
(tma_tensor_sfb.stride[0][0], ((x, x), x_times_3)),
tma_tensor_sfb.stride[1],
tma_tensor_sfb.stride[2],
)
tma_tensor_sfb_new_layout = cute.make_layout(
new_shape, stride=new_stride
)
tma_tensor_sfb = cute.make_tensor(
tma_tensor_sfb.iterator, tma_tensor_sfb_new_layout
)
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)
self.num_tma_load_bytes = (
a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
) * atom_thr_size
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
tma_op_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor,
epi_smem_layout,
self.epi_tile,
)
self.tile_sched_params, grid = self._compute_grid(
c_tensor,
self.cta_tile_shape_mnk,
self.cluster_shape_mn,
max_active_clusters,
self.swizzle_size,
self.raster_along_m,
)
self.buffer_align_bytes = 1024
self._define_shared_storage()
self.kernel(
tiled_mma,
tiled_mma_sfb,
tma_op_a,
tma_tensor_a,
tma_op_b,
tma_tensor_b,
tma_op_sfa,
tma_tensor_sfa,
tma_op_sfb,
tma_tensor_sfb,
tma_op_c,
tma_tensor_c,
self.cluster_layout_vmnk,
self.cluster_layout_sfb_vmnk,
self.a_smem_layout_staged,
self.b_smem_layout_staged,
self.sfa_smem_layout_staged,
self.sfb_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
self.tile_sched_params,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=1,
)
return
def _define_shared_storage(self):
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[
self.c_dtype,
cute.cosize(self.c_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sA: cute.struct.Align[
cute.struct.MemRange[
self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
sB: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfa_smem_layout_staged),
],
self.buffer_align_bytes,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfb_smem_layout_staged),
],
self.buffer_align_bytes,
]
self.shared_storage = SharedStorage
@cute.kernel
def kernel(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_op_a: cute.CopyAtom,
gA_mkl: cute.Tensor,
tma_op_b: cute.CopyAtom,
gB_nkl: cute.Tensor,
tma_op_sfa: cute.CopyAtom,
gSFA_mkl: cute.Tensor,
tma_op_sfb: cute.CopyAtom,
gSFB_nkl: cute.Tensor,
tma_op_c: cute.CopyAtom,
gC_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: cute.ComposedLayout,
epi_tile: cute.Tile,
tile_sched_params: utils.PersistentTileSchedulerParams,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
if warp_idx == self.tma_warp_id:
cpasync.prefetch_descriptor(tma_op_a)
cpasync.prefetch_descriptor(tma_op_b)
cpasync.prefetch_descriptor(tma_op_sfa)
cpasync.prefetch_descriptor(tma_op_sfb)
cpasync.prefetch_descriptor(tma_op_c)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
bidx, bidy, bidz = cute.arch.block_idx()
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
cta_rank_in_cluster
)
tidx, _, _ = cute.arch.thread_idx()
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_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_acc_consumer_threads = len(self.epilog_warp_id) * (
2 if use_2cta_instrs else 1
)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_acc_consumer_threads
)
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
)
pipeline_init_arrive(cluster_shape_mn=self.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 = 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
)
lA_mkl = cute.local_tile(
gA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
lB_nkl = cute.local_tile(
gB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
lSFA_mkl = cute.local_tile(
gSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
lSFB_nkl = cute.local_tile(
gSFB_nkl,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
lC_mnl = cute.local_tile(
gC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(lA_mkl, mode=[3])
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(lA_mkl)
tCgB = thr_mma.partition_B(lB_nkl)
tCgSFA = thr_mma.partition_A(lSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(lSFB_nkl)
tCgC = thr_mma.partition_C(lC_mnl)
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, tAgA = cpasync.tma_partition(
tma_op_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, tBgB = cpasync.tma_partition(
tma_op_b,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB, 0, 3),
)
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cpasync.tma_partition(
tma_op_sfa,
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA, 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, tBgSFB = cpasync.tma_partition(
tma_op_sfb,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB, 0, 3),
)
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
tRegA = tiled_mma.make_fragment_A(sA)
tRegB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
if cutlass.const_expr(self.overlapping_accum):
num_acc_stage_overlapped = 2
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, num_acc_stage_overlapped)
)
tCtAcc_fake = cute.make_tensor(
tCtAcc_fake.iterator,
cute.make_layout(
tCtAcc_fake.shape,
stride=(
tCtAcc_fake.stride[0],
tCtAcc_fake.stride[1],
tCtAcc_fake.stride[2],
(256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1],
),
),
)
else:
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
if warp_idx == self.tma_warp_id:
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
current_tile = tile_sched.initial_work_tile_info()
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
while current_tile.is_valid_tile:
tile_coord = current_tile.tile_idx
mma_coord_mnl = (
tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
tile_coord[1],
tile_coord[2],
)
tAgA_slice = tAgA[
(None, mma_coord_mnl[0], None, mma_coord_mnl[2])
]
tBgB_slice = tBgB[
(None, mma_coord_mnl[1], None, mma_coord_mnl[2])
]
tAgSFA_slice = tAgSFA[
(None, mma_coord_mnl[0], None, mma_coord_mnl[2])
]
slice_n = mma_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_coord_mnl[2])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
cute.copy(
tma_op_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=a_full_mcast_mask,
)
cute.copy(
tma_op_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_op_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=sfa_full_mcast_mask,
)
cute.copy(
tma_op_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=sfb_full_mcast_mask,
)
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
tile_sched.advance_to_next_work()
current_tile = tile_sched.get_current_work()
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols, dtype=self.sf_dtype
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
sfb_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
current_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
while current_tile.is_valid_tile:
tile_coord = current_tile.tile_idx
mma_coord_mnl = (
tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
tile_coord[1],
tile_coord[2],
)
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_producer_state.phase ^ 1
else:
acc_stage_index = acc_producer_state.index
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]
ab_consumer_state.reset_count()
acc_producer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
tCtSFB_mma = tCtSFB
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
offset = (
cutlass.Int32(2)
if mma_coord_mnl[1] % 2 == 1
else cutlass.Int32(0)
)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
offset = cutlass.Int32((mma_coord_mnl[1] % 2) * 2)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(k_tile_cnt):
if is_leader_cta:
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
s2t_stage_coord = (
None,
None,
None,
None,
ab_consumer_state.index,
)
tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
cute.copy(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t_staged,
tCtSFA_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t_staged,
tCtSFB_compact_s2t,
)
num_kblocks = cute.size(tRegA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (
None,
None,
kblock_idx,
ab_consumer_state.index,
)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(
tcgen05.Field.SFA,
tCtSFA[sf_kblock_coord].iterator,
)
tiled_mma.set(
tcgen05.Field.SFB,
tCtSFB_mma[sf_kblock_coord].iterator,
)
cute.gemm(
tiled_mma,
tCtAcc,
tRegA[kblock_coord],
tRegB[kblock_coord],
tCtAcc,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
tile_sched.advance_to_next_work()
current_tile = tile_sched.get_current_work()
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epilog_tidx = tidx
(
tiled_copy_t2r,
tT2R_tAcc_base,
tT2R_rAcc,
) = self.epilog_tmem_copy_and_partition(
epilog_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tT2R_rC = cute.make_rmem_tensor(tT2R_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tR2S_rC, tR2S_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tT2R_rC, epilog_tidx, sC
)
(
tma_op_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(
epilog_tidx, tma_op_c, tCgC, epi_tile, sC
)
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
)
current_tile = tile_sched.initial_work_tile_info()
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage,
producer_group=c_producer_group,
)
while current_tile.is_valid_tile:
tile_coord = current_tile.tile_idx
mma_coord_mnl = (
tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
tile_coord[1],
tile_coord[2],
)
bSG_gC = bSG_gC_partitioned[
(
None,
None,
None,
*mma_coord_mnl,
)
]
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
reverse_subtile = (
cutlass.Boolean(True)
if acc_stage_index == 0
else cutlass.Boolean(False)
)
else:
acc_stage_index = acc_consumer_state.index
tT2R_tAcc = tT2R_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tT2R_tAcc = cute.group_modes(tT2R_tAcc, 3, cute.rank(tT2R_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tT2R_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in cutlass.range(subtile_cnt):
real_subtile_idx = subtile_idx
if cutlass.const_expr(self.overlapping_accum):
if reverse_subtile:
real_subtile_idx = (
self.cta_tile_shape_mnk[1] // self.epi_tile_n
- 1
- subtile_idx
)
tT2R_tAcc_mn = tT2R_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tT2R_tAcc_mn, tT2R_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tT2R_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tR2S_rC.store(acc_vec)
c_buffer = (
num_prev_subtiles + real_subtile_idx
) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tR2S_rC,
tR2S_sC[(None, None, None, c_buffer)],
)
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.copy(
tma_op_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
if cutlass.const_expr(not self.overlapping_accum):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tile_sched.advance_to_next_work()
current_tile = tile_sched.get_current_work()
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
c_pipeline.producer_tail()
def mainloop_s2t_copy_and_partition(self, sSF: cute.Tensor, tSF: cute.Tensor):
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,
):
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tT2R_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
)
tT2R_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
tT2R_rAcc = cute.make_rmem_tensor(
tT2R_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
return tiled_copy_t2r, tT2R_tAcc, tT2R_rAcc
def epilog_smem_copy_and_partition(
self,
tiled_copy_t2r: cute.TiledCopy,
tT2R_rC: cute.Tensor,
tidx: cutlass.Int32,
sC: cute.Tensor,
):
copy_atom_r2s = sm100_utils.get_smem_store_op(
self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
)
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tR2S_sC = thr_copy_r2s.partition_D(sC)
tR2S_rC = tiled_copy_r2s.retile(tT2R_rC)
return tiled_copy_r2s, tR2S_rC, tR2S_sC
def epilog_gmem_copy_and_partition(
self,
tidx: cutlass.Int32,
atom: cute.CopyAtom,
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
sC: cute.Tensor,
):
gC_epi = cute.flat_divide(gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile)
tma_op_c = atom
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
bSG_sC, bSG_gC = cpasync.tma_partition(
tma_op_c,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC_for_tma_partition,
)
return tma_op_c, bSG_sC, bSG_gC
@staticmethod
def _compute_stages(
tiled_mma: cute.TiledMma,
mma_tiler_mnk,
a_dtype,
b_dtype,
epi_tile,
c_dtype,
c_layout,
sf_dtype,
sf_vec_size,
smem_capacity,
occupancy,
):
num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2
num_c_stage = 2
a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
tiled_mma, mma_tiler_mnk, a_dtype, 1
)
b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
tiled_mma, mma_tiler_mnk, b_dtype, 1
)
sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
tiled_mma, mma_tiler_mnk, sf_vec_size, 1
)
sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
tiled_mma, mma_tiler_mnk, sf_vec_size, 1
)
c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(
c_dtype, c_layout, epi_tile, 1
)
ab_bytes_per_stage = (
cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
+ cute.size_in_bytes(b_dtype, b_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
)
mbar_helpers_bytes = 1024
c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one)
c_bytes = c_bytes_per_stage * num_c_stage
num_ab_stage = (
smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
) // ab_bytes_per_stage
num_c_stage += (
smem_capacity
- occupancy * ab_bytes_per_stage * num_ab_stage
- occupancy * (mbar_helpers_bytes + c_bytes)
) // (occupancy * c_bytes_per_stage)
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
def _compute_grid(
c: cute.Tensor,
cta_tile_shape_mnk,
cluster_shape_mn,
max_active_clusters: cutlass.Constexpr,
swizzle_size: int,
raster_along_m: bool,
):
problem_shape_ntile_mnl = (
cute.ceil_div(c.shape[0], cta_tile_shape_mnk[0]),
cute.ceil_div(c.shape[1], cta_tile_shape_mnk[1]),
c.shape[2],
)
tile_sched_params = utils.PersistentTileSchedulerParams(
problem_shape_ntile_mnl,
(cluster_shape_mn[0], cluster_shape_mn[1], 1),
swizzle_size=swizzle_size,
raster_along_m=raster_along_m,
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
return tile_sched_params, grid
def select_config(m: int, n: int, k: int) -> KernelConfig:
if n == 7168 and k >= 16384:
return CONFIGS["n7168_k16384"]
if n == 4096 and k >= 7168:
return CONFIGS["n4096_k7168"]
if n == 7168 and k <= 2048:
return CONFIGS["n7168_k2048"]
return CONFIGS["default"]
_compiled_kernel_cache = {}
def compile_kernel(cfg: KernelConfig):
compiled = _compiled_kernel_cache.get(cfg.name)
if compiled is not None:
return compiled
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
gemm = PersistentNvfp4Gemm(cfg)
hardware_info = utils.HardwareInfo()
max_active = hardware_info.get_max_active_clusters(
cfg.cluster_shape_mn[0] * cfg.cluster_shape_mn[1]
)
@cute.jit
def wrapper(
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
):
m, n, k, l = problem_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr,
cute.make_layout((cute.assume(m, 32), n, l), stride=(n, 1, m * n)),
)
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_tensor.shape, sf_vec_size
)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
gemm(
a_tensor,
b_tensor,
sfa_tensor,
sfb_tensor,
c_tensor,
max_active,
)
return
compiled = cute.compile(
wrapper,
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
(0, 0, 0, 0),
)
_compiled_kernel_cache[cfg.name] = compiled
return compiled
def custom_kernel(data: input_t) -> output_t:
a = data[0]
b = data[1]
sfa_permuted = data[4]
sfb_permuted = data[5]
c = data[6]
m, k, l = a.shape
n, _, _ = b.shape
k = k * 2
cfg = select_config(m, n, k)
try:
compiled_func = compile_kernel(cfg)
except Exception:
cfg = select_config(m, n, k)
compiled_func = compile_kernel(cfg)
a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
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
)
sfb_ptr = make_ptr(
sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
return cscrolls · 1396 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