v4.3 update. (#2709)

* v4.3 update.

* Update the cute_dsl_api changelog's doc link

* Update version to 4.3.0

* Update the example link

* Update doc to encourage user to install DSL from requirements.txt

---------

Co-authored-by: Larry Wu <larwu@nvidia.com>
This commit is contained in:
Junkai-Wu
2025-10-21 14:26:30 -04:00
committed by GitHub
co-authored by Larry Wu
parent e6e2cc29f5
commit b1d6e2c9b3
244 changed files with 59272 additions and 10455 deletions
+102 -167
View File
@@ -38,6 +38,7 @@ import cutlass
import cutlass.cute as cute
import cutlass.cute.testing as testing
import cutlass.utils as utils
import cutlass.pipeline as pipeline
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.torch as cutlass_torch
@@ -152,12 +153,24 @@ class GroupedGemmKernel:
self.threads_per_cta = 32 * len(
(self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
)
# Set barrier id for cta sync, epilog sync, tmem ptr sync and tensormap update sync
self.cta_sync_bar_id = 0
self.epilog_sync_bar_id = 1
self.tmem_ptr_sync_bar_id = 2
# Barrier ID used by MMA/TMA warps to signal A/B tensormap initialization completion
self.tensormap_ab_init_bar_id = 4
# Set barrier for cta sync, epilog sync, tmem ptr sync and tensormap update sync
self.cta_sync_barrier = pipeline.NamedBarrier(
barrier_id=1,
num_threads=self.threads_per_cta,
)
self.epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=2,
num_threads=32 * len(self.epilog_warp_id),
)
self.tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=3,
num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
)
# Barrier used by MMA/TMA warps to signal A/B tensormap initialization completion
self.tensormap_ab_init_barrier = pipeline.NamedBarrier(
barrier_id=4,
num_threads=32 * (len(self.epilog_warp_id) + 1),
)
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
self.num_tma_load_bytes = 0
@@ -251,14 +264,6 @@ class GroupedGemmKernel:
self.num_epi_stage,
)
tensor_smem_bytes = self._get_tensor_smem_bytes(
self.a_smem_layout_staged,
self.a_dtype,
self.b_smem_layout_staged,
self.b_dtype,
self.epi_smem_layout_staged,
self.c_dtype,
)
mbar_smem_bytes = self._get_mbar_smem_bytes(
num_acc_stage=self.num_acc_stage,
num_ab_stage=self.num_ab_stage,
@@ -390,15 +395,12 @@ class GroupedGemmKernel:
# Setup TMA store for C
tma_atom_c = None
tma_tensor_c = None
c_cta_v_layout = cute.composition(
cute.make_identity_layout(initial_c.shape), self.epi_tile
)
epi_smem_layout = cute.slice_(self.epi_smem_layout_staged, (None, None, 0))
tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
initial_c,
epi_smem_layout,
c_cta_v_layout,
self.epi_tile,
)
self.tile_sched_params, grid = self._compute_grid(
@@ -558,8 +560,6 @@ class GroupedGemmKernel:
ab_empty_mbar_ptr = storage.ab_empty_mbar_ptr.data_ptr()
acc_full_mbar_ptr = storage.acc_full_mbar_ptr.data_ptr()
acc_empty_mbar_ptr = storage.acc_empty_mbar_ptr.data_ptr()
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
# init barrier for loading A, B with TMA
if warp_idx == self.epilog_warp_id[0]:
@@ -579,13 +579,13 @@ class GroupedGemmKernel:
acc_empty_mbar_ptr + acc_stage, 8 if use_2cta_instrs else 4
)
# Tensor memory dealloc barrier init
if use_2cta_instrs:
if warp_idx == self.tma_warp_id:
num_tmem_dealloc_threads = 32
with cute.arch.elect_one():
cute.arch.mbarrier_init(
tmem_dealloc_mbar_ptr, num_tmem_dealloc_threads
)
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,
)
cute.arch.mbarrier_init_fence()
# Cluster arrive after barrier init
@@ -721,9 +721,7 @@ class GroupedGemmKernel:
if cute.size(self.cluster_shape_mn) > 1:
cute.arch.cluster_wait()
else:
cute.arch.barrier(
barrier_id=self.cta_sync_bar_id, number_of_threads=self.threads_per_cta
)
self.cta_sync_barrier.arrive_and_wait()
#
# Get tensormap buffer address
@@ -826,10 +824,7 @@ class GroupedGemmKernel:
# wait tensormap initialization complete before update
if tensormap_init_done == False:
if cutlass.const_expr(self.delegate_tensormap_ab_init):
cute.arch.barrier(
barrier_id=self.tensormap_ab_init_bar_id,
number_of_threads=64,
)
self.tensormap_ab_init_barrier.arrive_and_wait()
tensormap_manager.fence_tensormap_initialization()
tensormap_init_done = True
@@ -951,33 +946,13 @@ class GroupedGemmKernel:
# Specialized MMA warp
#
if warp_idx == self.mma_warp_id:
# initialize tensormap A, B for TMA warp
if cutlass.const_expr(self.delegate_tensormap_ab_init):
tensormap_manager.init_tensormap_from_atom(
tma_atom_a, tensormap_a_init_ptr, self.mma_warp_id
)
tensormap_manager.init_tensormap_from_atom(
tma_atom_b, tensormap_b_init_ptr, self.mma_warp_id
)
# signal tensormap initialization has finished
cute.arch.barrier(
barrier_id=self.tensormap_ab_init_bar_id, number_of_threads=64
)
# Bar sync for retrieve tmem ptr from shared mem
tmem_ptr_read_threads = 32 * len((self.mma_warp_id, *self.epilog_warp_id))
cute.arch.barrier(
barrier_id=self.tmem_ptr_sync_bar_id,
number_of_threads=tmem_ptr_read_threads,
)
tmem.wait_for_alloc()
#
# Retrieving tensor memory ptr and make accumulator tensor
#
tmem_ptr = cute.arch.retrieve_tmem_ptr(
self.acc_dtype,
alignment=16,
ptr_to_buffer_holding_addr=tmem_holding_buf,
)
tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout)
@@ -1019,22 +994,22 @@ class GroupedGemmKernel:
# Peek (try_wait) AB buffer full for k_tile = 0
mma_rd_k_tile = cutlass.Int32(0)
smem_rd_buffer = (num_prev_k_blk + mma_rd_k_tile) % self.num_ab_stage
need_check_rd_buffer_full = (
mma_rd_k_tile < cur_k_tile_cnt and is_leader_cta
)
mma_rd_ab_full_phase = (
(num_prev_k_blk + mma_rd_k_tile) // self.num_ab_stage % 2
)
peek_ab_full_status = cute.arch.mbarrier_conditional_try_wait(
need_check_rd_buffer_full,
ab_full_mbar_ptr + smem_rd_buffer,
mma_rd_ab_full_phase,
)
#
# Wait for accumulator buffer empty
#
if is_leader_cta:
need_check_rd_buffer_full = (
mma_rd_k_tile < cur_k_tile_cnt and is_leader_cta
)
mma_rd_ab_full_phase = (
(num_prev_k_blk + mma_rd_k_tile) // self.num_ab_stage % 2
)
peek_ab_full_status = cute.arch.mbarrier_conditional_try_wait(
need_check_rd_buffer_full,
ab_full_mbar_ptr + smem_rd_buffer,
mma_rd_ab_full_phase,
)
#
# Wait for accumulator buffer empty
#
acc_empty_phase = (
tile_sched.num_tiles_executed // self.num_acc_stage % 2 ^ 1
)
@@ -1042,25 +1017,24 @@ class GroupedGemmKernel:
acc_empty_mbar_ptr + acc_buf_idx, acc_empty_phase
)
#
# Reset the ACCUMULATE field for each tile
#
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
#
# Reset the ACCUMULATE field for each tile
#
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
#
# Mma mainloop
#
for k_tile in range(cur_k_tile_cnt):
mma_rd_k_tile_next = cutlass.Int32(k_tile + 1)
smem_rd_buffer_next = (
num_prev_k_blk + mma_rd_k_tile_next
) % self.num_ab_stage
mma_rd_ab_full_phase_next = (
mma_rd_ab_full_phase ^ 1
if smem_rd_buffer_next == 0
else mma_rd_ab_full_phase
)
if is_leader_cta:
#
# Mma mainloop
#
for k_tile in range(cur_k_tile_cnt):
mma_rd_k_tile_next = cutlass.Int32(k_tile + 1)
smem_rd_buffer_next = (
num_prev_k_blk + mma_rd_k_tile_next
) % self.num_ab_stage
mma_rd_ab_full_phase_next = (
mma_rd_ab_full_phase ^ 1
if smem_rd_buffer_next == 0
else mma_rd_ab_full_phase
)
# Wait for AB buffer full
if peek_ab_full_status == 0:
cute.arch.mbarrier_wait(
@@ -1090,25 +1064,24 @@ class GroupedGemmKernel:
self.cta_group,
)
# Peek (try_wait) AB buffer full for k_tile = k_tile + 1
need_check_rd_buffer_full = (
mma_rd_k_tile_next < cur_k_tile_cnt and is_leader_cta
)
# Peek (try_wait) AB buffer full for k_tile = k_tile + 1
need_check_rd_buffer_full = (
mma_rd_k_tile_next < cur_k_tile_cnt and is_leader_cta
)
peek_ab_full_status = cute.arch.mbarrier_conditional_try_wait(
need_check_rd_buffer_full,
ab_full_mbar_ptr + smem_rd_buffer_next,
mma_rd_ab_full_phase_next,
)
peek_ab_full_status = cute.arch.mbarrier_conditional_try_wait(
need_check_rd_buffer_full,
ab_full_mbar_ptr + smem_rd_buffer_next,
mma_rd_ab_full_phase_next,
)
mma_rd_k_tile = mma_rd_k_tile_next
smem_rd_buffer = smem_rd_buffer_next
mma_rd_ab_full_phase = mma_rd_ab_full_phase_next
mma_rd_k_tile = mma_rd_k_tile_next
smem_rd_buffer = smem_rd_buffer_next
mma_rd_ab_full_phase = mma_rd_ab_full_phase_next
#
# Async arrive accumulator buffer full
#
if is_leader_cta:
#
# Async arrive accumulator buffer full
#
with cute.arch.elect_one():
tcgen05.commit(
acc_full_mbar_ptr + acc_buf_idx,
@@ -1126,6 +1099,16 @@ class GroupedGemmKernel:
# Specialized epilogue warps
#
if warp_idx < self.mma_warp_id:
# initialize tensormap A, B for TMA warp
if cutlass.const_expr(self.delegate_tensormap_ab_init):
tensormap_manager.init_tensormap_from_atom(
tma_atom_a, tensormap_a_init_ptr, self.epilog_warp_id[0]
)
tensormap_manager.init_tensormap_from_atom(
tma_atom_b, tensormap_b_init_ptr, self.epilog_warp_id[0]
)
# signal tensormap initialization has finished
self.tensormap_ab_init_barrier.arrive_and_wait()
# initialize tensorap for C
tensormap_manager.init_tensormap_from_atom(
tma_atom_c,
@@ -1133,30 +1116,17 @@ class GroupedGemmKernel:
self.epilog_warp_id[0],
)
# Alloc tensor memory buffer
if warp_idx == self.epilog_warp_id[0]:
cute.arch.alloc_tmem(
self.num_tmem_alloc_cols,
tmem_holding_buf,
is_two_cta=use_2cta_instrs,
)
tmem.allocate(self.num_tmem_alloc_cols)
#
# Bar sync for retrieve tensor memory ptr from shared memory
#
tmem_ptr_read_threads = 32 * len((self.mma_warp_id, *self.epilog_warp_id))
cute.arch.barrier(
barrier_id=self.tmem_ptr_sync_bar_id,
number_of_threads=tmem_ptr_read_threads,
)
tmem.wait_for_alloc()
#
# Retrieving tensor memory ptr and make accumulator tensor
#
tmem_ptr = cute.arch.retrieve_tmem_ptr(
self.acc_dtype,
alignment=16,
ptr_to_buffer_holding_addr=tmem_holding_buf,
)
tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout)
@@ -1172,7 +1142,7 @@ class GroupedGemmKernel:
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_fragment(tTR_rAcc.shape, self.c_dtype)
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
)
@@ -1303,11 +1273,7 @@ class GroupedGemmKernel:
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
epilog_threads = 32 * len(self.epilog_warp_id)
cute.arch.barrier(
barrier_id=self.epilog_sync_bar_id,
number_of_threads=epilog_threads,
)
self.epilog_sync_barrier.arrive_and_wait()
#
# store C to global memory with TMA
#
@@ -1325,10 +1291,7 @@ class GroupedGemmKernel:
cute.arch.cp_async_bulk_wait_group(
self.num_epi_stage - 1, read=True
)
cute.arch.barrier(
barrier_id=self.epilog_sync_bar_id,
number_of_threads=epilog_threads,
)
self.epilog_sync_barrier.arrive_and_wait()
#
# Async arrive accumulator buffer empty
#
@@ -1348,21 +1311,9 @@ class GroupedGemmKernel:
#
# Dealloc the tensor memory buffer
#
if warp_idx == self.epilog_warp_id[0]:
cute.arch.relinquish_tmem_alloc_permit(is_two_cta=use_2cta_instrs)
epilog_threads = 32 * len(self.epilog_warp_id)
cute.arch.barrier(
barrier_id=self.epilog_sync_bar_id, number_of_threads=epilog_threads
)
if warp_idx == self.epilog_warp_id[0]:
if use_2cta_instrs:
cute.arch.mbarrier_arrive(
tmem_dealloc_mbar_ptr, cta_rank_in_cluster ^ 1
)
cute.arch.mbarrier_wait(tmem_dealloc_mbar_ptr, 0)
cute.arch.dealloc_tmem(
tmem_ptr, self.num_tmem_alloc_cols, is_two_cta=use_2cta_instrs
)
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(tmem_ptr)
#
# Wait a/b buffer empty
@@ -1417,7 +1368,7 @@ class GroupedGemmKernel:
)
strides_tensor_gmem = strides_abc[(group_idx, tensor_index, None)]
strides_tensor_reg = cute.make_fragment(
strides_tensor_reg = cute.make_rmem_tensor(
cute.make_layout(2),
strides_abc.element_type,
)
@@ -1507,7 +1458,7 @@ class GroupedGemmKernel:
# (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_fragment(
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
@@ -1761,23 +1712,6 @@ class GroupedGemmKernel:
else:
raise ValueError(f"Invalid tensormap update mode: {tensormap_update_mode}")
@staticmethod
def _get_tensor_smem_bytes(
a_smem_layout_staged: cute.Layout,
a_dtype: Type[cutlass.Numeric],
b_smem_layout_staged: cute.Layout,
b_dtype: Type[cutlass.Numeric],
epi_smem_layout_staged: cute.Layout,
c_dtype: Type[cutlass.Numeric],
) -> int:
"""Compute the total SMEM consumption for tensor A, B and C."""
ab_bytes = cute.size_in_bytes(
a_dtype, a_smem_layout_staged
) + cute.size_in_bytes(b_dtype, b_smem_layout_staged)
epi_bytes = cute.size_in_bytes(c_dtype, epi_smem_layout_staged)
return ab_bytes + epi_bytes
@staticmethod
def _compute_num_tmem_alloc_cols(
tiled_mma: cute.TiledMma,
@@ -1821,7 +1755,7 @@ def create_tensor_and_stride(
is_dynamic_layout: bool = True,
torch_tensor_cpu: torch.Tensor = None,
) -> tuple[int, torch.Tensor, cute.Tensor, torch.Tensor, tuple[int, int]]:
"""Create a GPU tensor from scratch or based on an existing CPU tensor.
"""Create GPU tensor from either a new or existing CPU tensor.
:param torch_tensor_cpu: Optional existing CPU tensor to reuse. If None, creates a new one.
:type torch_tensor_cpu: torch.Tensor, optional
@@ -1967,7 +1901,7 @@ def run(
:return: Execution time of the GEMM kernel in microseconds
:rtype: float
"""
print(f"Running Blackwell Grouped GEMM test with:")
print("Running Blackwell Grouped GEMM test with:")
print(f"{num_groups} groups")
for i, (m, n, k, l) in enumerate(problem_sizes_mnkl):
print(f"Group {i}: {m}x{n}x{k}x{l}")
@@ -2102,6 +2036,7 @@ def run(
is_dynamic_layout=False,
assumed_align=16,
)
# layout (num_groups, 3, 2):(6, 2, 1)
tensor_of_strides_abc, tensor_of_strides_abc_torch = cutlass_torch.cute_tensor_like(
torch.tensor(strides_abc, dtype=torch.int32),