submission 506472
guaguabear · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3306 lines, June 9 Researcher Reciprocity License v1.0.
try.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-506472?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:8a69bce434f810a7652ac3f4def97b2026132c1dc4884ed884dfa61888fb206b
license declaredunknown
license concludedunknown
authorsguaguabear
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled group GEMM.fused-epilogue
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(persistent-kernel
) -> tuple[utils.PersistentTileSchedulerParams, tuple[int, int, int]]:shared-memory
reserved_smem_bytes = 1024tcgen05
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONEwarp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
try.py3306 lines
import argparse
import functools
import numpy as np
import torch
from typing import List, Type, Tuple, Union
from task import input_t, output_t
from itertools import accumulate
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
import cutlass.utils as utils
import cutlass.pipeline as pipeline
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
from cutlass.cute.runtime import from_dlpack, make_ptr
# Pure Human, No AI !!!!
# Current Score (us):
# 31.0
# 20.6
# 8.8
# 6.5
# -----
# Several theoretically sound strategies were implemented but did not yield performance gains in this specific competition context:
#
# 1. TMA Reduction Operations
# Attempted to use `CopyReduceBulkTensorTileS2GOp` to mitigate wave quantization losses.
# However, the overhead of reduction management outweighed the utilization gains for these specific tile sizes.
#
# 2. Group Sorting for Memory Balancing
# Attempted to sort groups to balance memory access patterns.
# However, since all groups shared the same K dimension, the load remained naturally balanced as long as the tile distribution among CTAs was uniform.
# -----
# Brief Opt Log (bottom up)
# 13.8 | Sort groups by runtime work estimate (= group_num_clusters) to reduce tail effects.
# 13.8 | Set max_active_sms for better schedule (on P1). Finetune tile_mn, cluster_mn based on the flattened version
# 14.5 | Instead of dynamic TMA updates, pass all TMA atoms and tensors as direct inputs to minimize synchronization overhead.
# | The code is much more verbose now, but the performance gain is significant. It's especially effective for small group counts, including G=8.
# 15.4 | Finetune tile_mn, cluster_mn and tma cache_policy
# 15.9 | hand craft worktile_scheduler for g8 without using official StaticPersistentGroupTileScheduler.
# | reorder TileSchedule for active ctas with Block-wise distribution instead of Round-robin distribution
# | to reduce group change for each cta.
# |
# | e.g. 5 ctas, group_cnt = 5, tiles_per_gemm = 5,
# | (gx_y represents yth tile of xth gemm)
# | Round-robin distribution:
# | cta1: g1_1->g2_1->g3_1->g4_1->g5_1 (frequent group changes)
# | cta2: g1_2->g2_2->g3_2->g4_2->g5_2
# | cta3: g1_3->g2_3->g3_3->g4_3->g5_3
# | cta4: g1_4->g2_4->g3_4->g4_4->g5_4
# | cta5: g1_5->g2_5->g3_5->g4_5->g5_5
# |
# | Block-wise distribution:
# | cta1: g1_1->g1_2->g1_3->g1_4->g1_5 (few group changes -> less tma update)
# | cta2: g2_1->g2_2->g2_3->g2_4->g2_5
# | cta3: g3_1->g3_2->g3_3->g3_4->g3_5
# | cta4: g4_1->g4_2->g4_3->g4_4->g4_5
# | cta5: g5_1->g5_2->g5_3->g5_4->g5_5
# 16.5 | As group_cnt is static, 'group_search_result' can be calculated by linear searching
# | a prefix_group_sum_clusters (prepared during runtime on host side)
# 19.7 | instead of store ptrs in global mem, directly feed all ptrs into kernel (reasonable for small G)
# 23 | impl specialized kernel for G == 2 (as G is static, based on normal gemm example code)
# 31 | cache group_count relevant buffers to avoid repeatedly re-creattion or re-allocattion
# 200 | using https://github.com/NVIDIA/cutlass/blob/main/examples/python/CuTeDSL/blackwell/grouped_blockscaled_gemm.py as backbone
_TMA_CACHE_EVICT_NORMAL = 0x1000000000000000
_TMA_CACHE_EVICT_FIRST = 0x12F0000000000000
_TMA_CACHE_EVICT_LAST = 0x14F0000000000000
USE_UPDATE_DEVICE_TMA_DESCRIPTOR = False
# cache_key(k): tile_mn | cluster_mn | cache_policy | max_active_sms
config_map = {
7168: ((256, 256), (2, 1), _TMA_CACHE_EVICT_FIRST, 128),
2048: ((128, 128), (1, 1), _TMA_CACHE_EVICT_LAST, 148),
4096: ((256, 128), (2, 1), _TMA_CACHE_EVICT_FIRST, 148),
1536: ((128, 128), (1, 1), _TMA_CACHE_EVICT_FIRST, 148),
}
debug_map = {
7168: True,
2048: False,
4096: False,
1536: False,
}
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
# Please ensure rows and cols are multiples of 128 and 4 respectively
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
# Pad the input matrix if necessary
if padded_rows != rows or padded_cols != cols:
padded = torch.nn.functional.pad(
input_matrix,
(0, padded_cols - cols, 0, padded_rows - rows),
mode="constant",
value=0,
)
else:
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def ref_kernel(data):
"""
PyTorch reference implementation of NVFP4 block-scaled group GEMM.
"""
abc_tensors, sfasfb_tensors, _, problem_sizes = data
result_tensors = []
for i, (
(a_ref, b_ref, c_ref),
(sfa_ref, sfb_ref),
(m, n, k, l),
) in enumerate(
zip(
abc_tensors,
sfasfb_tensors,
problem_sizes,
)
):
for l_idx in range(l):
# Convert the scale factor tensor to blocked format
scale_a = to_blocked(sfa_ref[:, :, l_idx])
scale_b = to_blocked(sfb_ref[:, :, l_idx])
# (m, k) @ (n, k).T -> (m, n)
res = torch._scaled_mm(
a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2),
b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2),
scale_a.cuda(),
scale_b.cuda(),
bias=None,
out_dtype=torch.float16,
)
c_ref[:, :, l_idx] = res
result_tensors.append((c_ref))
return result_tensors
def compute_cluster_info(
MList,
n,
mma_tiler_n,
cluster_shape_mn,
):
cta_tile_shape_mn = [128, mma_tiler_n]
cluster_tile_shape_mn = tuple(x * y for x, y in zip(cta_tile_shape_mn, cluster_shape_mn))
total_num_clusters = 0
group_num_clusters = []
for m in MList:
num_clusters_mn = tuple(
(x + y - 1) // y for x, y in zip((m, n), cluster_tile_shape_mn)
)
tmp = functools.reduce(lambda x, y: x * y, num_clusters_mn)
total_num_clusters += tmp
group_num_clusters.append(tmp)
return total_num_clusters, group_num_clusters
class GroupGemm:
reserved_smem_bytes = 1024
bytes_per_tensormap = 128
num_tensormaps = 5
tensor_memory_management_bytes = 12
def __init__(
self,
cache_key,
n,
group_cnt,
):
mma_tiler_mn = config_map[cache_key][0]
cluster_shape_mn = config_map[cache_key][1]
self.debug = debug_map[cache_key]
max_active_sms = config_map[cache_key][3]
# G | K | N are static, hence can be cached
self.group_count = group_cnt
self.N = n
self.K = cache_key
self.host_interface = None
self.max_active_clusters = max_active_sms // (cluster_shape_mn[0] * cluster_shape_mn[1])
self.k_tile_cnt = cache_key // 256
# Only for TMA update device descriptor version
tensormap_shape = (
max_active_sms,
GroupGemm.num_tensormaps,
GroupGemm.bytes_per_tensormap // 8,
)
self.tensor_of_tensormap = torch.empty(tensormap_shape, dtype=torch.int64, device="cuda")
self.cute_ptr_of_tensor_of_tensormap = make_ptr(
cutlass.Int64,
self.tensor_of_tensormap.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
# ---
self.acc_dtype = cutlass.Float32
self.sf_vec_size = 16
self.use_2cta_instrs = mma_tiler_mn[0] == 256
self.cluster_shape_mn = cluster_shape_mn
self.mma_tiler = (*mma_tiler_mn, 1)
self.cta_dim_n = (self.N + mma_tiler_mn[1] - 1) // mma_tiler_mn[1]
self.cache_policy = config_map[cache_key][2]
self.ignore_pipeline_sync = (self.cluster_shape_mn[0] == 1 and self.cluster_shape_mn[1] == 1) or (self.use_2cta_instrs and self.cluster_shape_mn[0] == 2)
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.occupancy = 1
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)),
)
# Only for TMA update device descriptor version
self.tensormap_ab_init_barrier = pipeline.NamedBarrier(
barrier_id=3,
num_threads=64,
)
# ---
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
SM100_TMEM_CAPACITY_COLUMNS = 512
self.num_tmem_alloc_cols = SM100_TMEM_CAPACITY_COLUMNS
def _setup_attributes(self):
# Compute mma instruction shapes
# (MMA_Tile_Shape_M, MMA_Tile_Shape_N, MMA_Inst_Shape_K)
self.mma_inst_shape_mn = (
self.mma_tiler[0],
self.mma_tiler[1],
)
# (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K)
self.mma_inst_shape_mn_sfb = (
self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
cute.round_up(self.mma_inst_shape_mn[1], 128),
)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
# Compute mma/cluster/tile shapes
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_tile_shape_mnk = tuple(
x * y for x, y in zip(self.cta_tile_shape_mnk, (*self.cluster_shape_mn, 1))
)
# Compute cluster layout
self.cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma.thr_id.shape,),
)
self.cluster_layout_sfb_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma_sfb.thr_id.shape,),
)
# Compute number of multicast CTAs for A/B
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
# Compute epilogue subtile
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
self.cta_tile_shape_mnk,
self.use_2cta_instrs,
self.c_layout,
self.c_dtype,
)
self.epi_tile_n = cute.size(self.epi_tile[1])
# Setup A/B/C stage count in shared memory and ACC stage count in tensor memory
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.group_count,
)
# Compute A/B/SFA/SFB/C shared memory layout
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,
)
# Overlap and double buffer accumulator when num_acc_stage == 1 for cta_tile_n = 256 case
self.overlapping_accum = False if self.group_count == 2 else (self.num_acc_stage == 1)
# Compute number of TMEM columns for SFA/SFB/Accumulator
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
self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage if not self.overlapping_accum else self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols
# Only when overlapping_accum is enabled, we need to release accumulator buffer early in epilogue
self.iter_acc_early_release_in_epilogue = self.num_sf_tmem_cols // self.epi_tile_n
def __call__(
self,
a_ptrs: List[cute.Pointer],
b_ptrs: List[cute.Pointer],
c_ptrs: List[cute.Pointer],
sfa_ptrs: List[cute.Pointer],
sfb_ptrs: List[cute.Pointer],
MList: List[int],
):
self.host_interface(
a_ptrs,
b_ptrs,
c_ptrs,
sfa_ptrs,
sfb_ptrs,
MList,
self.cute_ptr_of_tensor_of_tensormap, # Only for TMA update device descriptor with G = 8
)
@cute.jit
def kernel_call_flattened(
self,
a_ptrs: List[cute.Pointer],
b_ptrs: List[cute.Pointer],
c_ptrs: List[cute.Pointer],
sfa_ptrs: List[cute.Pointer],
sfb_ptrs: List[cute.Pointer],
MList: List[int],
ptr_of_tensor_of_tensormap: cute.Pointer,
):
(
total_num_clusters,
group_num_clusters,
) = compute_cluster_info(MList, self.N, self.mma_tiler[1], self.cluster_shape_mn)
prefix_group_num_clusters = list(accumulate(group_num_clusters, initial=0))
devide_bidz = total_num_clusters % self.max_active_clusters
if devide_bidz == 0:
devide_bidz = self.max_active_clusters
self.a_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN
self.b_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN
self.sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN
self.c_dtype: Type[cutlass.Numeric] = cutlass.Float16
self.a_major_mode = cutlass.cute.nvgpu.tcgen05.OperandMajorMode.K
self.b_major_mode = cutlass.cute.nvgpu.tcgen05.OperandMajorMode.K
self.c_layout = cutlass.utils.LayoutEnum.ROW_MAJOR
self._setup_attributes()
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
# Setup TMA load for A
a_op = sm100_utils.cluster_shape_to_tma_atom_A(
self.cluster_shape_mn, tiled_mma.thr_id
)
a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
# Setup TMA load for B
b_op = sm100_utils.cluster_shape_to_tma_atom_B(
self.cluster_shape_mn, tiled_mma.thr_id
)
b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
# Setup TMA load for SFA
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfa_smem_layout = cute.slice_(
self.sfa_smem_layout_staged, (None, None, None, 0)
)
# Setup TMA load for SFB
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfb_smem_layout = cute.slice_(
self.sfb_smem_layout_staged, (None, None, None, 0)
)
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
# Setup TMA store for C
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
self.tile_sched_params, grid = self._compute_grid(
total_num_clusters, self.cluster_shape_mn, self.max_active_clusters
)
self.buffer_align_bytes = 1024
self.size_tensormap_in_i64 = (
GroupGemm.num_tensormaps
* GroupGemm.bytes_per_tensormap
// 8
)
a_tensors = [
cute.make_tensor(
a_ptrs[i],
cute.make_layout((MList[i], self.K, 1),
stride=(self.K, 1, MList[i] * self.K))
) for i in range(self.group_count)
]
b_tensors = [
cute.make_tensor(
b_ptrs[i],
cute.make_layout((self.N, self.K, 1),
stride=(self.K, 1, self.N * self.K))
) for i in range(self.group_count)
]
sfa_tensors = [
cute.make_tensor(
sfa_ptrs[i],
blockscaled_utils.tile_atom_to_shape_SF(
a_tensors[i].shape, self.sf_vec_size
)
) for i in range(self.group_count)
]
sfb_tensors = [
cute.make_tensor(
sfb_ptrs[i],
blockscaled_utils.tile_atom_to_shape_SF(
b_tensors[i].shape, self.sf_vec_size
)
) for i in range(self.group_count)
]
c_tensors = [
cute.make_tensor(
c_ptrs[i],
cute.make_layout((MList[i], self.N, 1),
stride=(self.N, 1, MList[i] * self.N))
) for i in range(self.group_count)
]
tma_atoms_tensors_a = [
cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensors[i],
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
) for i in range(self.group_count)
]
tma_atoms_tensors_b = [
cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensors[i],
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
) for i in range(self.group_count)
]
tma_atoms_tensors_sfa = [
cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensors[i],
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
) for i in range(self.group_count)
]
tma_atoms_tensors_sfb = [
cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensors[i],
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
) for i in range(self.group_count)
]
tma_atoms_tensors_c = [
cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensors[i],
epi_smem_layout,
self.epi_tile,
) for i in range(self.group_count)
]
# Define shared storage for kernel
@cute.struct
class SharedStorage:
tensormap_buffer: cute.struct.MemRange[
cutlass.Int64, self.size_tensormap_in_i64
]
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
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
cute.struct.MemRange[
self.c_dtype,
cute.cosize(self.c_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
# (MMA, MMA_M, MMA_K, STAGE)
sA: cute.struct.Align[
cute.struct.MemRange[
self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
# (MMA, MMA_N, MMA_K, STAGE)
sB: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
# (MMA, MMA_M, MMA_K, STAGE)
sSFA: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
],
self.buffer_align_bytes,
]
# (MMA, MMA_N, MMA_K, STAGE)
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
# Launch the kernel synchronously
self.kernel_flattened(
tiled_mma,
tiled_mma_sfb,
tma_atoms_tensors_a,
tma_atoms_tensors_b,
tma_atoms_tensors_sfa,
tma_atoms_tensors_sfb,
tma_atoms_tensors_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,
MList,
total_num_clusters,
prefix_group_num_clusters,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=1,
)
return
@cute.kernel
def kernel_flattened(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atoms_tensors_a,
tma_atoms_tensors_b,
tma_atoms_tensors_sfa,
tma_atoms_tensors_sfb,
tma_atoms_tensors_c,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
MList,
total_num_clusters,
prefix_group_num_clusters,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# if warp_idx == self.tma_warp_id:
# for i in cutlass.range_constexpr(self.group_count):
# cpasync.prefetch_descriptor(tma_atoms_tensors_a[i][0])
# cpasync.prefetch_descriptor(tma_atoms_tensors_b[i][0])
# cpasync.prefetch_descriptor(tma_atoms_tensors_sfa[i][0])
# cpasync.prefetch_descriptor(tma_atoms_tensors_sfb[i][0])
# cpasync.prefetch_descriptor(tma_atoms_tensors_c[i][0])
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,
)
if cutlass.const_expr(not self.ignore_pipeline_sync):
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
)
# all
mA_mkl_all = [tensorA for (_, tensorA) in tma_atoms_tensors_a]
mB_nkl_all = [tensorB for (_, tensorB) in tma_atoms_tensors_b]
mSFA_mkl_all = [tensorSFA for (_, tensorSFA) in tma_atoms_tensors_sfa]
mSFB_nkl_all = [tensorSFB for (_, tensorSFB) in tma_atoms_tensors_sfb]
tma_atom_a_all = [atomA for (atomA, _) in tma_atoms_tensors_a]
tma_atom_b_all = [atomB for (atomB, _) in tma_atoms_tensors_b]
tma_atom_sfa_all = [atomSFA for (atomSFA, _) in tma_atoms_tensors_sfa]
tma_atom_sfb_all = [atomSFB for (atomSFB, _) in tma_atoms_tensors_sfb]
gA_mkl_all = [
cute.local_tile(
tensorA, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
) for tensorA in mA_mkl_all
]
gB_nkl_all = [
cute.local_tile(
tensorB, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
) for tensorB in mB_nkl_all
]
gSFA_mkl_all = [
cute.local_tile(
tensorSFA, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
) for tensorSFA in mSFA_mkl_all
]
gSFB_nkl_all = [
cute.local_tile(
tensorSFB, cute.slice_(self.mma_tiler_sfb, (0, None, None)), (None, None, None)
) for tensorSFB in mSFB_nkl_all
]
# ----
#
# Partition global tensor for TiledMMA_A/B/C
#
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA_all = [thr_mma.partition_A(gA) for gA in gA_mkl_all]
tCgB_all = [thr_mma.partition_B(gB) for gB in gB_nkl_all]
tCgSFA_all = [thr_mma.partition_A(gSFA) for gSFA in gSFA_mkl_all]
tCgSFB_all = [thr_mma_sfb.partition_B(gSFB) for gSFB in gSFB_nkl_all]
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, _ = cpasync.tma_partition(
tma_atom_a_all[0],
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA_all[0], 0, 3),
)
tAgA_all = [
(cpasync.tma_partition(
tma_atom_a_all[i],
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA_all[i], 0, 3),
))[1] for i in range(self.group_count)
]
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, _ = cpasync.tma_partition(
tma_atom_b_all[0],
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB_all[0], 0, 3),
)
tBgB_all = [
(cpasync.tma_partition(
tma_atom_b_all[i],
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB_all[i], 0, 3),
))[1] for i in range(self.group_count)
]
sfa_cta_layout = a_cta_layout
tAsSFA, _ = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfa_all[0],
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA_all[0], 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA_all = [
cute.filter_zeros((cute.nvgpu.cpasync.tma_partition(
tma_atom_sfa_all[i],
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA_all[i], 0, 3),
))[1]) for i in range(self.group_count)
]
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, _ = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb_all[0],
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB_all[0], 0, 3),
)
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB_all = [
cute.filter_zeros((cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb_all[i],
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB_all[i], 0, 3),
))[1]) for i in range(self.group_count)
]
#
# Partition shared/tensor memory tensor for TiledMMA_A/B/C
#
# (MMA, MMA_M, MMA_K, STAGE)
tCrA = tiled_mma.make_fragment_A(sA)
# (MMA, MMA_N, MMA_K, STAGE)
tCrB = tiled_mma.make_fragment_B(sB)
# (MMA, MMA_M, MMA_N)
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)
)
# (MMA, MMA_M, MMA_N, STAGE)
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:
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
#
# Cluster wait before tensor memory alloc
#
if cutlass.const_expr(not self.ignore_pipeline_sync):
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
#
# Specialized TMA load warp
#
if warp_idx == self.tma_warp_id:
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
for tilez_idx in cutlass.range(bidz, total_num_clusters, self.max_active_clusters):
# for tilez_idx in cutlass.range(lower_bound, upper_bound):
(
cur_group_idx,
_,
cta_tile_idx_m,
cta_tile_idx_n
) = self.dispatch_tile_info((bidx, bidy, tilez_idx), MList, prefix_group_num_clusters)
mma_tile_coord_mnl = (
cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
cta_tile_idx_n,
0,
)
if cur_group_idx == 0:
tAgA_slice = tAgA_all[0][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[0][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[0][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[0][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[0],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[0],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[0],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[0],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 1:
tAgA_slice = tAgA_all[1][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[1][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[1][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[1][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[1],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[1],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[1],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[1],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
if cutlass.const_expr(self.group_count == 8):
if cur_group_idx == 2:
tAgA_slice = tAgA_all[2][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[2][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[2][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[2][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[2],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[2],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[2],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[2],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 3:
tAgA_slice = tAgA_all[3][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[3][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[3][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[3][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[3],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[3],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[3],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[3],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 4:
tAgA_slice = tAgA_all[4][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[4][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[4][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[4][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[4],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[4],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[4],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[4],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 5:
tAgA_slice = tAgA_all[5][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[5][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[5][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[5][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[5],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[5],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[5],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[5],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 6:
tAgA_slice = tAgA_all[6][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[6][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[6][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[6][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[6],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[6],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[6],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[6],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
elif cur_group_idx == 7:
tAgA_slice = tAgA_all[7][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
tBgB_slice = tBgB_all[7][
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
tAgSFA_slice = tAgSFA_all[7][
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB_all[7][
(None, slice_n, None, mma_tile_coord_mnl[2])
]
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
#
# Tma load loop
#
for k_tile in cutlass.range(0, self.k_tile_cnt, 1, unroll=2):
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# TMA load A/B/SFA/SFB
cute.copy(
tma_atom_a_all[7],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_b_all[7],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfa_all[7],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
cute.copy(
tma_atom_sfb_all[7],
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,
cache_policy=cutlass.Int64(cutlass.Int64(self.cache_policy).ir_value()),
)
# Prefetch: Rolling prefetch for next tiles
# if self.prefetch_enabled:
# if k_tile < k_tile_cnt - self.prefetch_dist:
# future_k_tile = ab_producer_state.count + self.prefetch_dist
# cute.prefetch(
# tma_atom_a,
# tAgA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_b,
# tBgB_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfa,
# tAgSFA_slice[(None, future_k_tile)],
# )
# cute.prefetch(
# tma_atom_sfb,
# tBgSFB_slice[(None, future_k_tile)],
# )
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < self.k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
ab_pipeline.producer_tail(ab_producer_state)
#
# Specialized MMA warp
#
if warp_idx == self.mma_warp_id:
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + 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()
# )
# work_tile = tile_sched.initial_work_tile_info()
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
# while work_tile.is_valid_tile:
# cur_tile_coord = work_tile.tile_idx
# for tilez_idx in cutlass.range(lower_bound, upper_bound):
for tilez_idx in cutlass.range(bidz, total_num_clusters, self.max_active_clusters):
(
cur_group_idx,
problem_shape_m,
cta_tile_idx_m,
cta_tile_idx_n
) = self.dispatch_tile_info((bidx, bidy, tilez_idx), MList, prefix_group_num_clusters)
mma_tile_coord_mnl = (
cta_tile_idx_m // cute.size(tiled_mma.thr_id.shape),
cta_tile_idx_n,
0,
)
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()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < self.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):
# If this is an ODD tile, shift the TMEM start address for cta_tile_shape_n=192 case by two words (ignores first 64 columns of SFB)
offset = cutlass.Int32(2) if mma_tile_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):
# Move in increments of 64 columns of SFB
offset = cutlass.Int32((mma_tile_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)
#
# Reset the ACCUMULATE field for each tile
#
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
#
# Mma mainloop
#
for k_tile in range(self.k_tile_cnt):
if is_leader_cta:
# Conditionally wait for AB buffer full
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
# Copy SFA/SFB from smem to tmem
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,
)
# tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB
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,
)
# Set SFA/SFB tensor to tiled_mma
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(
tcgen05.Field.SFA,
tCtSFA[sf_kblock_coord].iterator,
)
tiled_mma.set(
tcgen05.Field.SFB,
tCtSFB_mma[sf_kblock_coord].iterator,
)
cute.gemm(
tiled_mma,
tCtAcc,
tCrA[kblock_coord],
tCrB[kblock_coord],
tCtAcc,
)
# Enable accumulate on tCtAcc after first kblock
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
# Async arrive AB buffer empty
ab_pipeline.consumer_release(ab_consumer_state)
# Peek (try_wait) AB buffer full for k_tile = k_tile + 1
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < self.k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
#
# Async arrive accumulator buffer full
#
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
#
# Advance to next tile
#
# tile_sched.advance_to_next_work()
# work_tile = tile_sched.get_current_work()
#
# Wait for accumulator buffer empty
#
if cutlass.const_expr(not self.group_count == 2):
acc_pipeline.producer_tail(acc_producer_state)
#
# Specialized epilogue warps
#
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
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,
)
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
mC_mnl_all = [tensorC for (_, tensorC) in tma_atoms_tensors_c]
tma_atom_c_all = [atomC for (atomC, _) in tma_atoms_tensors_c]
gC_mnl_all = [
cute.local_tile(
tensorC, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
) for tensorC in mC_mnl_all
]
tCgC_all = [thr_mma.partition_C(gC) for gC in gC_mnl_all]
epi_tidx = tidx
tiled_copy_t2r, tTR_tAcc_base, tTR_rAcc = (
self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC_all[0], epi_tile, use_2cta_instrs
)
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
bSG_sC, bSG_gC_partitioned_all = (
self.epilog_gmem_copy_and_partition_all(
epi_tidx, tma_atom_c_all, tCgC_all, epi_tile, sC
)
)
# for tilez_idx in cutlass.range(lower_bound, upper_bound):
for tilez_idx in cutlass.range(bidz, total_num_clusters, self.max_active_clusters):
(
cur_group_idx,
problem_shape_m,
cta_tile_idx_m,
cta_tile_idx_n
) = self.dispatch_tile_info((bidx, bidy, tilez_idx), MList, prefix_group_num_clusters)
mma_tile_coord_mnl = (
cta_tile_idx_m
// cute.size(tiled_mma.thr_id.shape),
cta_tile_idx_n,
0,
)
if cur_group_idx == 0:
bSG_gC = bSG_gC_partitioned_all[0][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[0],
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):
if cutlass.const_expr(not self.group_count == 2):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
elif cur_group_idx == 1:
bSG_gC = bSG_gC_partitioned_all[1][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[1],
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):
if cutlass.const_expr(not self.group_count == 2):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
if cutlass.const_expr(self.group_count == 8):
if cur_group_idx == 2:
bSG_gC = bSG_gC_partitioned_all[2][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# if(self.debug):
# cute.printf(subtile_cnt)
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[2],
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()
elif cur_group_idx == 3:
bSG_gC = bSG_gC_partitioned_all[3][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[3],
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()
elif cur_group_idx == 4:
bSG_gC = bSG_gC_partitioned_all[4][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[4],
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()
elif cur_group_idx == 5:
bSG_gC = bSG_gC_partitioned_all[5][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[5],
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()
elif cur_group_idx == 6:
bSG_gC = bSG_gC_partitioned_all[6][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[6],
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()
elif cur_group_idx == 7:
bSG_gC = bSG_gC_partitioned_all[7][
(
None,
None,
None,
*mma_tile_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
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
# num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
num_prev_subtiles = 0
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
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
with cute.arch.elect_one():
cute.arch.fence_view_async_tmem_load()
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = acc_vec.to(self.c_dtype)
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
cute.arch.fence_proxy("async.shared", space="cta")
cute.copy(
tma_atom_c_all[7],
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()
tmem.free(acc_tmem_ptr)
# c_pipeline.producer_tail()
@cute.jit
def dispatch_tile_info(
self,
tile_idx: tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32],
MList,
prefix_group_num_clusters,
) -> tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32]:
m, n, l = tile_idx
group_idx = cutlass.Int32(-1)
problem_shape_m = cutlass.Int32(-1)
cta_tile_idx_m = cutlass.Int32(-1)
cta_tile_idx_n = cutlass.Int32(-1)
if cutlass.const_expr(self.group_count == 2):
if l < prefix_group_num_clusters[1]:
group_idx = 0
problem_shape_m = MList[0]
cta_tile_idx_m = l // self.cta_dim_n
cta_tile_idx_n = l % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[2]:
group_idx = 1
problem_shape_m = MList[1]
cta_tile_idx_m = (l - prefix_group_num_clusters[1]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[1]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
else:
if l < prefix_group_num_clusters[1]:
group_idx = 0
problem_shape_m = MList[0]
cta_tile_idx_m = (l - prefix_group_num_clusters[0]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[0]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[2]:
group_idx = 1
problem_shape_m = MList[1]
cta_tile_idx_m = (l - prefix_group_num_clusters[1]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[1]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[3]:
group_idx = 2
problem_shape_m = MList[2]
cta_tile_idx_m = (l - prefix_group_num_clusters[2]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[2]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[4]:
group_idx = 3
problem_shape_m = MList[3]
cta_tile_idx_m = (l - prefix_group_num_clusters[3]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[3]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[5]:
group_idx = 4
problem_shape_m = MList[4]
cta_tile_idx_m = (l - prefix_group_num_clusters[4]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[4]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[6]:
group_idx = 5
problem_shape_m = MList[5]
cta_tile_idx_m = (l - prefix_group_num_clusters[5]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[5]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
elif l < prefix_group_num_clusters[7]:
group_idx = 6
problem_shape_m = MList[6]
cta_tile_idx_m = (l - prefix_group_num_clusters[6]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[6]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
else:
group_idx = 7
problem_shape_m = MList[7]
cta_tile_idx_m = (l - prefix_group_num_clusters[7]) // self.cta_dim_n
cta_tile_idx_n = (l - prefix_group_num_clusters[7]) % self.cta_dim_n
cta_tile_idx_m = cta_tile_idx_m * self.cluster_shape_mn[0] + m
cta_tile_idx_n = cta_tile_idx_n * self.cluster_shape_mn[1] + n
return (group_idx, problem_shape_m, cta_tile_idx_m, cta_tile_idx_n)
@cute.jit
def make_tensor_abc_for_tensormap_update(
self,
ptrs: List[cute.Pointer],
group_idx: cutlass.Int32,
problem_shape_m: cutlass.Int32,
tensor_index: int,
):
tensor_gmem_ptr = ptrs[0]
for i in cutlass.range_constexpr(1, self.group_count):
if group_idx == i:
tensor_gmem_ptr = ptrs[i]
if cutlass.const_expr(tensor_index == 0): # tensor A
m = problem_shape_m
k = self.K
return cute.make_tensor(
tensor_gmem_ptr,
cute.make_layout((m, k, 1), stride=(k, 1, 0)),
)
elif cutlass.const_expr(tensor_index == 1): # tensor B
n = self.N
k = self.K
return cute.make_tensor(
tensor_gmem_ptr,
cute.make_layout((n, k, 1), stride=(k, 1, 0)),
)
else: # tensor C
m = problem_shape_m
n = self.N
return cute.make_tensor(
tensor_gmem_ptr,
cute.make_layout((m, n, 1), stride=(n, 1, 0)),
)
@cute.jit
def make_tensor_sfasfb_for_tensormap_update(
self,
ptrs: List[cute.Pointer],
group_idx: cutlass.Int32,
problem_shape_m: cutlass.Int32,
tensor_index: int,
):
tensor_gmem_ptr = ptrs[0]
for i in cutlass.range_constexpr(1, self.group_count):
if group_idx == i:
tensor_gmem_ptr = ptrs[i]
if cutlass.const_expr(tensor_index == 0): # tensor SFA
m = problem_shape_m
k = self.K
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
(m, k, 1), self.sf_vec_size
)
return cute.make_tensor(
tensor_gmem_ptr,
sfa_layout,
)
else: # tensor SFB
n = self.N
k = self.K
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
(n, k, 1), self.sf_vec_size
)
return cute.make_tensor(
tensor_gmem_ptr,
sfb_layout,
)
def mainloop_s2t_copy_and_partition(
self,
sSF: cute.Tensor,
tSF: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
# (MMA, MMA_MN, MMA_K, STAGE)
tCsSF_compact = cute.filter_zeros(sSF)
# (MMA, MMA_MN, MMA_K)
tCtSF_compact = cute.filter_zeros(tSF)
# Make S2T CopyAtom and tiledCopy
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)
# ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
# ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE)
tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t, tCsSF_compact_s2t_
)
# ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K)
tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t
def epilog_tmem_copy_and_partition(
self,
tidx: cutlass.Int32,
tAcc: cute.Tensor,
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
use_2cta_instrs: Union[cutlass.Boolean, bool],
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
# Make tiledCopy for tensor memory load
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE)
tAcc_epi = cute.flat_divide(
tAcc[((None, None), 0, 0, None)],
epi_tile,
)
# (EPI_TILE_M, EPI_TILE_N)
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)
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC_mnl_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
# (T2R, T2R_M, T2R_N)
tTR_rAcc = cute.make_rmem_tensor(
tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
return tiled_copy_t2r, tTR_tAcc, tTR_rAcc
def epilog_smem_copy_and_partition(
self,
tiled_copy_t2r: cute.TiledCopy,
tTR_rC: cute.Tensor,
tidx: cutlass.Int32,
sC: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
copy_atom_r2s = sm100_utils.get_smem_store_op(
self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
)
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
# (R2S, R2S_M, R2S_N, PIPE_D)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
# (R2S, R2S_M, R2S_N)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(
self,
tidx: cutlass.Int32,
atom: Union[cute.CopyAtom, cute.TiledCopy],
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tma_atom_c = atom
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
# ((ATOM_V, REST_V), EPI_M, EPI_N)
# ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
bSG_sC, bSG_gC = cpasync.tma_partition(
tma_atom_c,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC_for_tma_partition,
)
return tma_atom_c, bSG_sC, bSG_gC
def epilog_gmem_copy_and_partition2(
self,
tidx: cutlass.Int32,
atom0: Union[cute.CopyAtom, cute.TiledCopy],
atom1: Union[cute.CopyAtom, cute.TiledCopy],
gC0_mnl: cute.Tensor,
gC1_mnl: cute.Tensor,
epi_tile: cute.Tile,
sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC0_epi = cute.flat_divide(
gC0_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
gC1_epi = cute.flat_divide(
gC1_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tma_atom_c0 = atom0
tma_atom_c1 = atom1
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
gC0_for_tma_partition = cute.group_modes(gC0_epi, 0, 2)
gC1_for_tma_partition = cute.group_modes(gC1_epi, 0, 2)
# ((ATOM_V, REST_V), EPI_M, EPI_N)
# ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
bSG_sC, bSG_gC0 = cpasync.tma_partition(
tma_atom_c0,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC0_for_tma_partition,
)
_, bSG_gC1 = cpasync.tma_partition(
tma_atom_c1,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC1_for_tma_partition,
)
return tma_atom_c0, tma_atom_c1, bSG_sC, bSG_gC0, bSG_gC1
def epilog_gmem_copy_and_partition_all(
self,
tidx: cutlass.Int32,
atom_all,
gC_mnl_all,
epi_tile: cute.Tile,
sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL)
gC_epi_all = [
cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
) for gC_mnl in gC_mnl_all
]
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
bSG_SC_gC_all = [
cpasync.tma_partition(
atom_all[i],
0,
cute.make_layout(1),
sC_for_tma_partition,
cute.group_modes(gC_epi_all[i], 0, 2),
) for i in range(self.group_count)
]
bSG_sC = bSG_SC_gC_all[0][0]
bSG_gC_all = [bSG_gC for (_, bSG_gC) in bSG_SC_gC_all]
return bSG_sC, bSG_gC_all
@staticmethod
def _compute_stages(
tiled_mma: cute.TiledMma,
mma_tiler_mnk: Tuple[int, int, int],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
epi_tile: cute.Tile,
c_dtype: Type[cutlass.Numeric],
c_layout: utils.LayoutEnum,
sf_dtype: Type[cutlass.Numeric],
sf_vec_size: int,
smem_capacity: int,
occupancy: int,
group_count: int,
) -> Tuple[int, int, int]:
# ACC stages
num_acc_stage = 1 if (group_count == 2 or mma_tiler_mnk[1] == 256) else 2
# Default C stages
num_c_stage = 2
# Calculate smem layout and size for one stage of A, B, SFA, SFB and C
a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
tiled_mma,
mma_tiler_mnk,
a_dtype,
1, # a tmp 1 stage is provided
)
b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
tiled_mma,
mma_tiler_mnk,
b_dtype,
1, # a tmp 1 stage is provided
)
sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
1, # a tmp 1 stage is provided
)
sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
1, # a tmp 1 stage is provided
)
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
# Calculate A/B/SFA/SFB stages:
# Start with total smem per CTA (capacity / occupancy)
# Subtract reserved bytes and initial C stages bytes
# Divide remaining by bytes needed per A/B/SFA/SFB stage
num_ab_stage = (
smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
) // ab_bytes_per_stage
# Refine epilogue stages:
# Calculate remaining smem after allocating for A/B/SFA/SFB stages and reserved bytes
# Add remaining unused smem to epilogue
num_c_stage += (
smem_capacity
- occupancy * ab_bytes_per_stage * num_ab_stage
- occupancy * (mbar_helpers_bytes + c_bytes)
) // (occupancy * c_bytes_per_stage)
# print(num_acc_stage, num_ab_stage, num_c_stage)
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
def _compute_grid(
total_num_clusters: int,
cluster_shape_mn: tuple[int, int],
max_active_clusters: cutlass.Constexpr[int],
) -> tuple[utils.PersistentTileSchedulerParams, tuple[int, int, int]]:
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
)
return tile_sched_params, grid
_compiled_kernel_cache = {}
def compile_kernel(cache_key, n, group_cnt):
global _compiled_kernel_cache
if cache_key in _compiled_kernel_cache:
return _compiled_kernel_cache[cache_key]
a_ptrs = [
make_ptr(
cutlass.Float4E2M1FN,
0, # fake ptrs for compile
cute.AddressSpace.gmem,
assumed_align=16,
)
for i in range(group_cnt)
]
b_ptrs = [
make_ptr(
cutlass.Float4E2M1FN,
0, # fake ptrs for compile
cute.AddressSpace.gmem,
assumed_align=16,
)
for i in range(group_cnt)
]
c_ptrs = [
make_ptr(
cutlass.Float16,
0, # fake ptrs for compile
cute.AddressSpace.gmem,
assumed_align=16,
)
for i in range(group_cnt)
]
sfa_ptrs = [
make_ptr(
cutlass.Float8E4M3FN,
0, # fake ptrs for compile
cute.AddressSpace.gmem,
assumed_align=16,
)
for i in range(group_cnt)
]
sfb_ptrs = [
make_ptr(
cutlass.Float8E4M3FN,
0, # fake ptrs for compile
cute.AddressSpace.gmem,
assumed_align=16,
)
for i in range(group_cnt)
]
# M is dynamic, fake MList for compile
MList = [0 for i in range(group_cnt)]
# K | N | G is static, which can be cached
group_gemm = GroupGemm(
cache_key,
n,
group_cnt,
)
group_gemm.host_interface = cute.compile(
group_gemm.kernel_call_flattened,
a_ptrs,
b_ptrs,
c_ptrs,
sfa_ptrs,
sfb_ptrs,
MList,
group_gemm.cute_ptr_of_tensor_of_tensormap,
)
_compiled_kernel_cache[cache_key] = group_gemm
return group_gemm
def custom_kernel(data):
abc_tensors, _, sfasfb_reordered_tensors, problem_sizes_mnkl = data
# take k for cache_key, as k is static during runtime
cache_key = problem_sizes_mnkl[0][2] # k
# N and G are static, can be cached when compiling
n = problem_sizes_mnkl[0][1]
group_cnt = len(problem_sizes_mnkl)
# optimize the shapes counted in leaderboard
if cache_key in [7168, 2048, 4096, 1536]:
group_gemm = compile_kernel(cache_key, n, group_cnt)
# M is dynamic, should be fed into kernel during runtime.
# Bucket by runtime work estimate (= group_num_clusters), then stable order within bucket.
# Restore output order after kernel launch.
MList = [m for (m, _, _, _) in problem_sizes_mnkl]
_, group_num_clusters = compute_cluster_info(
MList,
n,
group_gemm.mma_tiler[1],
group_gemm.cluster_shape_mn,
)
# buckets: key=workload bucket (group_num_clusters), value=list of original indices (stable)
buckets = {}
for i, w in enumerate(group_num_clusters):
buckets.setdefault(w, []).append(i)
# Heavy bucket first, stable within each bucket
sort_idx = []
for w in sorted(buckets.keys(), reverse=True):
sort_idx.extend(buckets[w])
if any(sort_idx[i] != i for i in range(group_cnt)):
abc_tensors_sorted = [abc_tensors[i] for i in sort_idx]
sfasfb_reordered_tensors_sorted = [sfasfb_reordered_tensors[i] for i in sort_idx]
problem_sizes_mnkl_sorted = [problem_sizes_mnkl[i] for i in sort_idx]
else:
abc_tensors_sorted = abc_tensors
sfasfb_reordered_tensors_sorted = sfasfb_reordered_tensors
problem_sizes_mnkl_sorted = problem_sizes_mnkl
a_ptrs = [
make_ptr(
cutlass.Float4E2M1FN,
abc_t[0].data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
for abc_t in abc_tensors_sorted
]
b_ptrs = [
make_ptr(
cutlass.Float4E2M1FN,
abc_t[1].data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
for abc_t in abc_tensors_sorted
]
c_ptrs = [
make_ptr(
cutlass.Float16,
abc_t[2].data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
for abc_t in abc_tensors_sorted
]
sfa_ptrs = [
make_ptr(
cutlass.Float8E4M3FN,
sfab_t[0].data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
for sfab_t in sfasfb_reordered_tensors_sorted
]
sfb_ptrs = [
make_ptr(
cutlass.Float8E4M3FN,
sfab_t[1].data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
for sfab_t in sfasfb_reordered_tensors_sorted
]
MList_sorted = [m for (m, _, _, _) in problem_sizes_mnkl_sorted]
group_gemm(a_ptrs, b_ptrs, c_ptrs, sfa_ptrs, sfb_ptrs, MList_sorted)
return [c for (_, _, c) in abc_tensors]
else:
return ref_kernel(data)
scrolls · 3306 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 506313.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON