v4.4 tag release update. (#3032)

This commit is contained in:
Junkai-Wu
2026-02-14 12:27:58 +08:00
committed by GitHub
parent 01687cfba1
commit d4bbf728ca
140 changed files with 41624 additions and 3691 deletions

View File

@@ -335,9 +335,11 @@ class GroupedGemmKernel:
:type stream: cuda.CUstream
:raises TypeError: If A and B data types do not match.
"""
self.a_dtype = initial_a.element_type
self.b_dtype = initial_b.element_type
self.c_dtype = initial_c.element_type
self.a_major_mode = utils.LayoutEnum.from_tensor(initial_a).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(initial_b).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(initial_c)
@@ -475,6 +477,7 @@ class GroupedGemmKernel:
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
stream=stream,
min_blocks_per_mp=1,
)
return
@@ -553,28 +556,38 @@ class GroupedGemmKernel:
tensormap_c_smem_ptr = (
tensormap_b_smem_ptr + GroupedGemmKernel.bytes_per_tensormap // 8
)
ab_full_mbar_ptr = storage.ab_full_mbar_ptr.data_ptr()
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()
# init barrier for loading A, B with TMA
if warp_idx == self.epilog_warp_id[0]:
for k_stage in range(self.num_ab_stage):
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
with cute.arch.elect_one():
cute.arch.mbarrier_init(ab_full_mbar_ptr + k_stage, 1)
cute.arch.mbarrier_init(
ab_empty_mbar_ptr + k_stage, num_tma_producer
)
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,
)
# Accumulator barrier init
if warp_idx == self.mma_warp_id:
for acc_stage in range(self.num_acc_stage):
with cute.arch.elect_one():
cute.arch.mbarrier_init(acc_full_mbar_ptr + acc_stage, 1)
cute.arch.mbarrier_init(
acc_empty_mbar_ptr + acc_stage, 8 if use_2cta_instrs else 4
)
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,
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
@@ -747,10 +760,26 @@ class GroupedGemmKernel:
tensormap_b_init_ptr = tensormap_b_ptr
tensormap_c_init_ptr = tensormap_c_ptr
#
# Persistent tile scheduling loop
#
# When the problem shapes are on device, we launch one CTA per SM.
# The if condition later prevents the warps from extra CTAs from doing any work.
tile_sched = utils.StaticPersistentGroupTileScheduler.create(
tile_sched_params,
bid,
grid_dim,
self.cluster_tile_shape_mnk,
utils.create_initial_search_state(),
group_count,
problem_sizes_mnkl,
)
initial_work_tile_info = tile_sched.initial_work_tile_info()
#
# Specialized TMA load warp
#
if warp_idx == self.tma_warp_id:
if warp_idx == self.tma_warp_id and initial_work_tile_info.is_valid_tile:
# Initialize tensormaps for A, B
if cutlass.const_expr(self.delegate_tensormap_ab_init == False):
tensormap_manager.init_tensormap_from_atom(
@@ -759,185 +788,161 @@ class GroupedGemmKernel:
tensormap_manager.init_tensormap_from_atom(
tma_atom_b, tensormap_b_init_ptr, self.tma_warp_id
)
#
# Persistent tile scheduling loop
#
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, bid, grid_dim
)
# grouped gemm tile scheduler helper will compute the group index for the tile we're working on
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count,
tile_sched_params,
self.cluster_tile_shape_mnk,
utils.create_initial_search_state(),
)
tensormap_init_done = cutlass.Boolean(False)
# tile count we have searched
total_k_tile_cnt = cutlass.Int32(0)
# group index of last tile
last_group_idx = cutlass.Int32(-1)
work_tile = tile_sched.initial_work_tile_info()
work_tile = initial_work_tile_info
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
grouped_gemm_cta_tile_info = group_gemm_ts_helper.delinearize_z(
cur_tile_coord,
problem_sizes_mnkl,
)
grouped_gemm_cta_tile_info = work_tile.group_search_result
cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
cur_group_idx = grouped_gemm_cta_tile_info.group_idx
is_group_changed = cur_group_idx != last_group_idx
# skip tensormap update if we're working on the same group
if is_group_changed:
real_tensor_a = self.make_tensor_for_tensormap_update(
cur_group_idx,
self.a_dtype,
(
grouped_gemm_cta_tile_info.problem_shape_m,
grouped_gemm_cta_tile_info.problem_shape_n,
grouped_gemm_cta_tile_info.problem_shape_k,
),
strides_abc,
ptrs_abc,
0, # 0 for tensor A
# Do not load any data if cur_k_tile_cnt is 0
if not is_k_tile_cnt_zero:
is_group_changed = cur_group_idx != last_group_idx
# skip tensormap update if we're working on the same group
if is_group_changed:
real_tensor_a = self.make_tensor_for_tensormap_update(
cur_group_idx,
self.a_dtype,
(
grouped_gemm_cta_tile_info.problem_shape_m,
grouped_gemm_cta_tile_info.problem_shape_n,
grouped_gemm_cta_tile_info.problem_shape_k,
),
strides_abc,
ptrs_abc,
0, # 0 for tensor A
)
real_tensor_b = self.make_tensor_for_tensormap_update(
cur_group_idx,
self.b_dtype,
(
grouped_gemm_cta_tile_info.problem_shape_m,
grouped_gemm_cta_tile_info.problem_shape_n,
grouped_gemm_cta_tile_info.problem_shape_k,
),
strides_abc,
ptrs_abc,
1, # 1 for tensor B
)
# wait tensormap initialization complete before update
if not tensormap_init_done:
if cutlass.const_expr(self.delegate_tensormap_ab_init):
self.tensormap_ab_init_barrier.arrive_and_wait()
tensormap_manager.fence_tensormap_initialization()
tensormap_init_done = True
tensormap_manager.update_tensormap(
(real_tensor_a, real_tensor_b),
(tma_atom_a, tma_atom_b),
(tensormap_a_ptr, tensormap_b_ptr),
self.tma_warp_id,
(tensormap_a_smem_ptr, tensormap_b_smem_ptr),
)
mma_tile_coord_mnl = (
grouped_gemm_cta_tile_info.cta_tile_idx_m
// cute.size(tiled_mma.thr_id.shape),
grouped_gemm_cta_tile_info.cta_tile_idx_n,
0,
)
real_tensor_b = self.make_tensor_for_tensormap_update(
cur_group_idx,
self.b_dtype,
(
grouped_gemm_cta_tile_info.problem_shape_m,
grouped_gemm_cta_tile_info.problem_shape_n,
grouped_gemm_cta_tile_info.problem_shape_k,
),
strides_abc,
ptrs_abc,
1, # 1 for tensor B
)
# wait tensormap initialization complete before update
if tensormap_init_done == False:
#
# Slice to per mma tile index
#
# ((atom_v, rest_v), RestK)
tAgA_slice = tAgA[
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
# ((atom_v, rest_v), RestK)
tBgB_slice = tBgB[
(None, mma_tile_coord_mnl[1], 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 < cur_k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
# ensure the update to tensormap has completed before using it
if is_group_changed:
tensormap_manager.fence_tensormap_update(tensormap_a_ptr)
tensormap_manager.fence_tensormap_update(tensormap_b_ptr)
#
# Tma load loop
#
for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
# Wait for AB buffer empty
ab_pipeline.producer_acquire(
ab_producer_state, peek_ab_empty_status
)
# Load A/B with TMA
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=a_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(
tensormap_a_ptr,
cute.AddressSpace.generic,
),
)
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(
ab_producer_state
),
mcast_mask=b_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(
tensormap_b_ptr,
cute.AddressSpace.generic,
),
)
# 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 < cur_k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
else:
# If tensormap initialization is not done, wait for it to complete
if not tensormap_init_done:
if cutlass.const_expr(self.delegate_tensormap_ab_init):
self.tensormap_ab_init_barrier.arrive_and_wait()
tensormap_manager.fence_tensormap_initialization()
tensormap_init_done = True
tensormap_manager.update_tensormap(
(real_tensor_a, real_tensor_b),
(tma_atom_a, tma_atom_b),
(tensormap_a_ptr, tensormap_b_ptr),
self.tma_warp_id,
(tensormap_a_smem_ptr, tensormap_b_smem_ptr),
)
mma_tile_coord_mnl = (
grouped_gemm_cta_tile_info.cta_tile_idx_m
// cute.size(tiled_mma.thr_id.shape),
grouped_gemm_cta_tile_info.cta_tile_idx_n,
0,
)
#
# Slice to per mma tile index
#
# ((atom_v, rest_v), RestK)
tAgA_slice = tAgA[
(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
]
# ((atom_v, rest_v), RestK)
tBgB_slice = tBgB[
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
]
num_prev_k_blk = total_k_tile_cnt
total_k_tile_cnt += cur_k_tile_cnt
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
tma_wr_k_tile = cutlass.Int32(0)
smem_wr_buffer = (num_prev_k_blk + tma_wr_k_tile) % self.num_ab_stage
tma_wr_ab_empty_phase = (
num_prev_k_blk + tma_wr_k_tile
) // self.num_ab_stage % 2 ^ 1
peek_ab_empty_status = cute.arch.mbarrier_conditional_try_wait(
tma_wr_k_tile < cur_k_tile_cnt,
ab_empty_mbar_ptr + smem_wr_buffer,
tma_wr_ab_empty_phase,
)
# ensure the update to tensormap has completed before using it
if is_group_changed:
tensormap_manager.fence_tensormap_update(tensormap_a_ptr)
tensormap_manager.fence_tensormap_update(tensormap_b_ptr)
#
# Tma load loop
#
for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
tma_wr_k_tile_next = tma_wr_k_tile + 1
smem_wr_buffer_next = (
num_prev_k_blk + tma_wr_k_tile_next
) % self.num_ab_stage
tma_wr_ab_empty_phase_next = (
tma_wr_ab_empty_phase ^ 1
if smem_wr_buffer_next == 0
else tma_wr_ab_empty_phase
)
smem_full_mbar_ptr = ab_full_mbar_ptr + smem_wr_buffer
# Wait for AB buffer empty
if peek_ab_empty_status == 0:
cute.arch.mbarrier_wait(
ab_empty_mbar_ptr + smem_wr_buffer, tma_wr_ab_empty_phase
)
# Arrive AB buffer and expect full transaction bytes
if is_leader_cta:
with cute.arch.elect_one():
cute.arch.mbarrier_arrive_and_expect_tx(
smem_full_mbar_ptr, self.num_tma_load_bytes
)
# Load A/B with TMA
cute.copy(
tma_atom_a,
tAgA_slice[(None, tma_wr_k_tile)],
tAsA[(None, smem_wr_buffer)],
tma_bar_ptr=smem_full_mbar_ptr,
mcast_mask=a_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(
tensormap_a_ptr,
cute.AddressSpace.generic,
),
)
cute.copy(
tma_atom_b,
tBgB_slice[(None, tma_wr_k_tile)],
tBsB[(None, smem_wr_buffer)],
tma_bar_ptr=smem_full_mbar_ptr,
mcast_mask=b_full_mcast_mask,
tma_desc_ptr=tensormap_manager.get_tensormap_ptr(
tensormap_b_ptr,
cute.AddressSpace.generic,
),
)
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
peek_ab_empty_status = cute.arch.mbarrier_conditional_try_wait(
tma_wr_k_tile_next < cur_k_tile_cnt,
ab_empty_mbar_ptr + smem_wr_buffer_next,
tma_wr_ab_empty_phase_next,
)
tma_wr_k_tile = tma_wr_k_tile_next
smem_wr_buffer = smem_wr_buffer_next
tma_wr_ab_empty_phase = tma_wr_ab_empty_phase_next
# Advance to next tile
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
last_group_idx = cur_group_idx
#
# Wait A/B buffer empty
#
ab_pipeline.producer_tail(ab_producer_state)
#
# Specialized MMA warp
#
if warp_idx == self.mma_warp_id:
if warp_idx == self.mma_warp_id and initial_work_tile_info.is_valid_tile:
# Bar sync for retrieve tmem ptr from shared mem
tmem.wait_for_alloc()
@@ -951,63 +956,42 @@ class GroupedGemmKernel:
#
# Persistent tile scheduling loop
#
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, bid, grid_dim
work_tile = initial_work_tile_info
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
# grouped gemm tile scheduler helper will compute the group index for the tile we're working on
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count,
tile_sched_params,
self.cluster_tile_shape_mnk,
utils.create_initial_search_state(),
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
work_tile = tile_sched.initial_work_tile_info()
# tile count we have searched
total_k_tile_cnt = cutlass.Int32(0)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
# MMA warp is only interested in number of tiles along K dimension
(
cur_k_tile_cnt,
cur_group_idx,
) = group_gemm_ts_helper.search_cluster_tile_count_k(
cur_tile_coord,
problem_sizes_mnkl,
)
# Set tensor memory buffer for current tile
acc_buf_idx = tile_sched.num_tiles_executed % self.num_acc_stage
# (MMA, MMA_M, MMA_N)
tCtAcc = tCtAcc_base[(None, None, None, acc_buf_idx)]
cur_group_idx = work_tile.group_search_result.group_idx
problem_shape_k = work_tile.group_search_result.problem_shape_k
num_prev_k_blk = total_k_tile_cnt
total_k_tile_cnt += cur_k_tile_cnt
# MMA warp is only interested in number of tiles along K dimension
cur_k_tile_cnt = (
problem_shape_k + self.cluster_tile_shape_mnk[2] - 1
) // self.cluster_tile_shape_mnk[2]
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
# (MMA, MMA_M, MMA_N)
tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
# 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
ab_consumer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
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,
)
if ab_consumer_state.count < cur_k_tile_cnt:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
#
# Wait for accumulator buffer empty
#
acc_empty_phase = (
tile_sched.num_tiles_executed // self.num_acc_stage % 2 ^ 1
)
cute.arch.mbarrier_wait(
acc_empty_mbar_ptr + acc_buf_idx, acc_empty_phase
)
if not is_k_tile_cnt_zero:
acc_pipeline.producer_acquire(acc_producer_state)
#
# Reset the ACCUMULATE field for each tile
@@ -1017,26 +1001,20 @@ class GroupedGemmKernel:
#
# 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
)
for k_tile in cutlass.range(0, cur_k_tile_cnt, 1, unroll=1):
# Wait for AB buffer full
if peek_ab_full_status == 0:
cute.arch.mbarrier_wait(
ab_full_mbar_ptr + smem_rd_buffer, mma_rd_ab_full_phase
)
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
# tCtAcc += tCrA * tCrB
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, smem_rd_buffer)
kblock_coord = (
None,
None,
kblock_idx,
ab_consumer_state.index,
)
cute.gemm(
tiled_mma,
@@ -1049,48 +1027,37 @@ class GroupedGemmKernel:
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
# Async arrive AB buffer empty
with cute.arch.elect_one():
tcgen05.commit(
ab_empty_mbar_ptr + smem_rd_buffer,
ab_empty_mcast_mask,
self.cta_group,
)
ab_pipeline.consumer_release(ab_consumer_state)
# 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,
)
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
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < cur_k_tile_cnt:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
#
# Async arrive accumulator buffer full
#
with cute.arch.elect_one():
tcgen05.commit(
acc_full_mbar_ptr + acc_buf_idx,
acc_full_mcast_mask,
self.cta_group,
)
if not is_k_tile_cnt_zero:
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
#
acc_pipeline.producer_tail(acc_producer_state)
#
# Specialized epilogue warps
#
if warp_idx < self.mma_warp_id:
if warp_idx < self.mma_warp_id and initial_work_tile_info.is_valid_tile:
# initialize tensormap A, B for TMA warp
if cutlass.const_expr(self.delegate_tensormap_ab_init):
tensormap_manager.init_tensormap_from_atom(
@@ -1147,32 +1114,32 @@ class GroupedGemmKernel:
#
# Persistent tile scheduling loop
#
tile_sched = utils.StaticPersistentTileScheduler.create(
tile_sched_params, bid, grid_dim
)
# grouped gemm tile scheduler helper will compute the group index for the tile we're working on
group_gemm_ts_helper = utils.GroupedGemmTileSchedulerHelper(
group_count,
tile_sched_params,
self.cluster_tile_shape_mnk,
utils.create_initial_search_state(),
)
work_tile = tile_sched.initial_work_tile_info()
work_tile = initial_work_tile_info
# wait tensormap initialization complete before update
tensormap_manager.fence_tensormap_initialization()
# tile count we have searched
total_k_tile_cnt = cutlass.Int32(0)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
# Threads/warps participating in tma store pipeline
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_epi_stage,
producer_group=c_producer_group,
)
# group index of last tile
last_group_idx = cutlass.Int32(-1)
while work_tile.is_valid_tile:
cur_tile_coord = work_tile.tile_idx
grouped_gemm_cta_tile_info = group_gemm_ts_helper.delinearize_z(
cur_tile_coord,
problem_sizes_mnkl,
)
grouped_gemm_cta_tile_info = work_tile.group_search_result
cur_group_idx = grouped_gemm_cta_tile_info.group_idx
cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
is_k_tile_cnt_zero = cur_k_tile_cnt == 0
is_group_changed = cur_group_idx != last_group_idx
# We still need to store 0s when k_tile_cnt is 0
if is_group_changed:
# construct tensor C based on real address, shape and stride information
real_tensor_c = self.make_tensor_for_tensormap_update(
@@ -1201,8 +1168,6 @@ class GroupedGemmKernel:
grouped_gemm_cta_tile_info.cta_tile_idx_n,
0,
)
cur_k_tile_cnt = grouped_gemm_cta_tile_info.cta_tile_count_k
total_k_tile_cnt += cur_k_tile_cnt
#
# Slice to per mma tile index
@@ -1216,17 +1181,16 @@ class GroupedGemmKernel:
*mma_tile_coord_mnl,
)
]
# Set tensor memory buffer for current tile
acc_buf_idx = tile_sched.num_tiles_executed % self.num_acc_stage
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_buf_idx)]
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_consumer_state.index)
]
#
# Wait for accumulator buffer full
#
acc_full_phase = tile_sched.num_tiles_executed // self.num_acc_stage % 2
cute.arch.mbarrier_wait(acc_full_mbar_ptr + acc_buf_idx, acc_full_phase)
# Not waiting for accumulator buffer full when k_tile_cnt is 0
if not is_k_tile_cnt_zero:
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))
@@ -1240,28 +1204,34 @@ class GroupedGemmKernel:
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in range(subtile_cnt):
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
#
# Convert to output type
#
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
tRS_rC.store(acc_vec.to(self.c_dtype))
#
# Store C to shared memory
#
epi_buffer = (num_prev_subtiles + subtile_idx) % self.num_epi_stage
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
if not is_k_tile_cnt_zero:
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
#
# Convert to output type
#
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
tRS_rC.store(acc_vec.to(self.c_dtype))
else:
tRS_rC.fill(0)
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, epi_buffer)],
)
# Fence and barrier to make sure shared memory store is visible to TMA store
cute.arch.fence_proxy("async.shared", space="cta")
cute.arch.fence_proxy(
"async.shared",
space="cta",
)
self.epilog_sync_barrier.arrive_and_wait()
#
# store C to global memory with TMA
@@ -1276,19 +1246,17 @@ class GroupedGemmKernel:
cute.AddressSpace.generic,
),
)
cute.arch.cp_async_bulk_commit_group()
cute.arch.cp_async_bulk_wait_group(
self.num_epi_stage - 1, read=True
)
# Fence and barrier to make sure shared memory store is visible to TMA store
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
#
# Async arrive accumulator buffer empty
#
with cute.arch.elect_one():
cute.arch.mbarrier_arrive(
acc_empty_mbar_ptr + acc_buf_idx,
cta_rank_in_cluster // 2 * 2 if use_2cta_instrs else None,
)
if not is_k_tile_cnt_zero:
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
#
# Advance to next tile
@@ -1305,13 +1273,9 @@ class GroupedGemmKernel:
tmem.free(tmem_ptr)
#
# Wait a/b buffer empty
# Wait for C store complete
#
if warp_idx == self.epilog_warp_id[0]:
cute.arch.mbarrier_wait(
(ab_empty_mbar_ptr + ((total_k_tile_cnt - 1) % self.num_ab_stage)),
(((total_k_tile_cnt - 1) // self.num_ab_stage) % 2),
)
c_pipeline.producer_tail()
@cute.jit
def make_tensor_for_tensormap_update(
@@ -1649,7 +1613,7 @@ class GroupedGemmKernel:
problem_shape_ntile_mnl, (*cluster_shape_mn, 1)
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
grid = utils.StaticPersistentGroupTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
@@ -1866,6 +1830,7 @@ def create_tensors_for_all_groups(
def run(
num_groups: int,
problem_sizes_mnkl: tuple[int, int, int, int],
host_problem_shape_available: bool,
ab_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
acc_dtype: Type[cutlass.Numeric],
@@ -1975,18 +1940,18 @@ def run(
c_major,
)
# Choose A, B, C with the smallest size to create initial tensormaps
key_size_a = lambda item: item[1][0] * item[1][2]
key_size_b = lambda item: item[1][1] * item[1][2]
key_size_c = lambda item: item[1][0] * item[1][1]
# Find the indices of the groups with the smallest tensor sizes
min_a_idx, _ = min(enumerate(problem_sizes_mnkl), key=key_size_a)
min_b_idx, _ = min(enumerate(problem_sizes_mnkl), key=key_size_b)
min_c_idx, _ = min(enumerate(problem_sizes_mnkl), key=key_size_c)
# Setup inital tensors for TMA of A,B and C
alignment = 16 # 16 bytes aligned
min_ab_size = alignment * 8 // ab_dtype.width
min_c_size = alignment * 8 // c_dtype.width
initial_cute_tensors_abc = [
cute_tensors_abc[min_a_idx][0], # A with smallest (m, k)
cute_tensors_abc[min_b_idx][1], # B with smallest (n, k)
cute_tensors_abc[min_c_idx][2], # C with smallest (m, n)
create_tensor_and_stride(1, min_ab_size, min_ab_size, a_major == "m", ab_dtype)[
2
],
create_tensor_and_stride(1, min_ab_size, min_ab_size, b_major == "n", ab_dtype)[
2
],
create_tensor_and_stride(1, min_c_size, min_c_size, c_major == "m", c_dtype)[2],
]
hardware_info = utils.HardwareInfo()
@@ -1994,6 +1959,7 @@ def run(
max_active_clusters = hardware_info.get_max_active_clusters(
cluster_shape_mn[0] * cluster_shape_mn[1]
)
# Prepare tensormap buffer for each SM
num_tensormap_buffers = sm_count
tensormap_shape = (
@@ -2069,9 +2035,19 @@ def run(
cluster_tile_shape_mn = compute_cluster_tile_shape(
mma_tiler_mn, cluster_shape_mn, use_2cta_instrs
)
total_num_clusters = compute_total_num_clusters(
problem_sizes_mnkl, cluster_tile_shape_mn
)
# If the host problem shape is available, we will launch the grid with only
# the necessary clusters. The function compute_total_num_clusters() does that.
# If the problem shape only exists on device, we will need to launch all active
# clusters possible on a device.
if host_problem_shape_available:
print("Problem shapes available on host and device")
total_num_clusters = compute_total_num_clusters(
problem_sizes_mnkl, cluster_tile_shape_mn
)
else:
print("Problem shapes available only on device")
total_num_clusters = max_active_clusters
# Initialize Stream
current_stream = cutlass_torch.default_stream()
@@ -2079,6 +2055,7 @@ def run(
# try to check CUDA version to decide the opt level
try:
from cutlass import CUDA_VERSION
opt_level = (
3
if CUDA_VERSION.major < 13
@@ -2131,6 +2108,9 @@ def run(
rtol=1e-05,
)
if iterations <= 0:
return 0
def generate_tensors():
# Reuse existing CPU tensors and create new GPU tensors from them
(
@@ -2150,9 +2130,15 @@ def run(
)
initial_cute_tensors_abc_workspace = [
cute_tensors_abc_workspace[min_a_idx][0], # A with smallest (m, k)
cute_tensors_abc_workspace[min_b_idx][1], # B with smallest (n, k)
cute_tensors_abc_workspace[min_c_idx][2], # C with smallest (m, n)
create_tensor_and_stride(
1, min_ab_size, min_ab_size, a_major == "m", ab_dtype
)[2],
create_tensor_and_stride(
1, min_ab_size, min_ab_size, b_major == "n", ab_dtype
)[2],
create_tensor_and_stride(
1, min_c_size, min_c_size, c_major == "m", c_dtype
)[2],
]
# Create new tensors for this workspace
@@ -2176,7 +2162,7 @@ def run(
is_dynamic_layout=False,
)
return testing.JitArguments(
args = testing.JitArguments(
initial_cute_tensors_abc_workspace[0],
initial_cute_tensors_abc_workspace[1],
initial_cute_tensors_abc_workspace[2],
@@ -2186,6 +2172,8 @@ def run(
tensormap_workspace,
current_stream,
)
args.add_to_scope([torch_tensors_abc_workspace])
return args
workspace_count = 1
if use_cold_l2:
@@ -2225,6 +2213,18 @@ def run(
iterations=iterations,
)
runtime_s = exec_time / 1.0e6
fmas = 0
for group in range(num_groups):
[M, N, K, _] = problem_sizes_mnkl[group]
fmas += M * N * K
flop = 2 * fmas
gflop = flop / 1.0e9
gflops = gflop / runtime_s
print("Average Runtime : ", exec_time / 1000, "ms")
print("GFLOPS : ", gflops)
return exec_time # Return execution time in microseconds
@@ -2270,15 +2270,20 @@ if __name__ == "__main__":
parser.add_argument(
"--num_groups",
type=int,
default=2,
default=3,
help="Number of groups",
)
parser.add_argument(
"--problem_sizes_mnkl",
type=parse_comma_separated_tuples,
default=((128, 128, 128, 1), (128, 128, 128, 1)),
default=((128, 128, 128, 1), (512, 128, 128, 1), (128, 256, 128, 1)),
help="a tuple of problem sizes for each group (comma-separated tuples)",
)
parser.add_argument(
"--host_problem_shape_available",
action="store_true",
help="Enable the compute of grid based upon host problem shape",
)
parser.add_argument(
"--mma_tiler_mn",
type=parse_comma_separated_ints,
@@ -2362,6 +2367,7 @@ if __name__ == "__main__":
run(
args.num_groups,
args.problem_sizes_mnkl,
args.host_problem_shape_available,
args.ab_dtype,
args.c_dtype,
args.acc_dtype,