v4.3 tag release update. (#2789)
This commit is contained in:
@@ -38,6 +38,7 @@ from cutlass.cute.nvgpu import cpasync, tcgen05
|
||||
import cutlass.torch as cutlass_torch
|
||||
import cutlass.utils as utils
|
||||
import cutlass.pipeline as pipeline
|
||||
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
|
||||
import cutlass.utils.blackwell_helpers as sm100_utils
|
||||
import cutlass.utils.blockscaled_layout as blockscaled_utils
|
||||
from cutlass.cute.runtime import from_dlpack
|
||||
@@ -208,17 +209,13 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.threads_per_cta = 32 * len(
|
||||
(self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
|
||||
)
|
||||
# Set barrier id for cta sync, epilogue sync and tmem ptr sync
|
||||
self.cta_sync_barrier = pipeline.NamedBarrier(
|
||||
barrier_id=1,
|
||||
num_threads=self.threads_per_cta,
|
||||
)
|
||||
# Set barrier id for epilogue sync and tmem ptr sync
|
||||
self.epilog_sync_barrier = pipeline.NamedBarrier(
|
||||
barrier_id=2,
|
||||
barrier_id=1,
|
||||
num_threads=32 * len(self.epilog_warp_id),
|
||||
)
|
||||
self.tmem_alloc_barrier = pipeline.NamedBarrier(
|
||||
barrier_id=3,
|
||||
barrier_id=2,
|
||||
num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
|
||||
)
|
||||
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
|
||||
@@ -288,6 +285,11 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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],
|
||||
)
|
||||
|
||||
# Compute cluster layout
|
||||
self.cluster_layout_vmnk = cute.tiled_divide(
|
||||
@@ -314,6 +316,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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(
|
||||
@@ -362,6 +365,19 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.num_c_stage,
|
||||
)
|
||||
|
||||
# Overlap and double buffer accumulator when num_acc_stage == 1 for cta_tile_n = 256 case
|
||||
self.overlapping_accum = 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
|
||||
|
||||
@cute.jit
|
||||
def __call__(
|
||||
self,
|
||||
@@ -640,6 +656,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
block=[self.threads_per_cta, 1, 1],
|
||||
cluster=(*self.cluster_shape_mn, 1),
|
||||
stream=stream,
|
||||
min_blocks_per_mp=1,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -726,6 +743,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
consumer_group=ab_pipeline_consumer_group,
|
||||
tx_count=self.num_tma_load_bytes,
|
||||
cta_layout_vmnk=cluster_layout_vmnk,
|
||||
defer_sync=True,
|
||||
)
|
||||
|
||||
# Initialize acc_pipeline (barrier) and states
|
||||
@@ -742,6 +760,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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
|
||||
@@ -754,8 +773,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
# Cluster arrive after barrier init
|
||||
if cute.size(self.cluster_shape_mn) > 1:
|
||||
cute.arch.cluster_arrive_relaxed()
|
||||
pipeline_init_arrive(cluster_shape_mn=self.cluster_shape_mn, is_relaxed=True)
|
||||
|
||||
#
|
||||
# Setup smem tensor A/B/SFA/SFB/C
|
||||
@@ -910,18 +928,34 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
tCrB = tiled_mma.make_fragment_B(sB)
|
||||
# (MMA, MMA_M, MMA_N)
|
||||
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
|
||||
# (MMA, MMA_M, MMA_N, STAGE)
|
||||
tCtAcc_fake = tiled_mma.make_fragment_C(
|
||||
cute.append(acc_shape, self.num_acc_stage)
|
||||
)
|
||||
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 cute.size(self.cluster_shape_mn) > 1:
|
||||
cute.arch.cluster_wait()
|
||||
else:
|
||||
self.cta_sync_barrier.arrive_and_wait()
|
||||
pipeline_init_wait(cluster_shape_mn=self.cluster_shape_mn)
|
||||
|
||||
#
|
||||
# Specialized TMA load warp
|
||||
@@ -1057,7 +1091,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
|
||||
# Make SFA tmem tensor
|
||||
sfa_tmem_ptr = cute.recast_ptr(
|
||||
acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base),
|
||||
acc_tmem_ptr + self.num_accumulator_tmem_cols,
|
||||
dtype=self.sf_dtype,
|
||||
)
|
||||
# (MMA, MMA_M, MMA_K)
|
||||
@@ -1071,9 +1105,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
|
||||
# Make SFB tmem tensor
|
||||
sfb_tmem_ptr = cute.recast_ptr(
|
||||
acc_tmem_ptr
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA),
|
||||
acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
|
||||
dtype=self.sf_dtype,
|
||||
)
|
||||
# (MMA, MMA_N, MMA_K)
|
||||
@@ -1122,9 +1154,15 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
cur_tile_coord[2],
|
||||
)
|
||||
|
||||
# Get accumulator stage index
|
||||
if cutlass.const_expr(self.overlapping_accum):
|
||||
acc_stage_index = acc_producer_state.phase ^ 1
|
||||
else:
|
||||
acc_stage_index = acc_producer_state.index
|
||||
|
||||
# Set tensor memory buffer for current tile
|
||||
# (MMA, MMA_M, MMA_N)
|
||||
tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)]
|
||||
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]
|
||||
|
||||
# Peek (try_wait) AB buffer full for k_tile = 0
|
||||
ab_consumer_state.reset_count()
|
||||
@@ -1146,8 +1184,8 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
offset = cutlass.Int32(2) if mma_tile_coord_mnl[1] % 2 == 1 else cutlass.Int32(0)
|
||||
shifted_ptr = cute.recast_ptr(
|
||||
acc_tmem_ptr
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
|
||||
+ self.num_accumulator_tmem_cols
|
||||
+ self.num_sfa_tmem_cols
|
||||
+ offset,
|
||||
dtype=self.sf_dtype,
|
||||
)
|
||||
@@ -1156,9 +1194,9 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
# 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
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
|
||||
acc_tmem_ptr
|
||||
+ self.num_accumulator_tmem_cols
|
||||
+ self.num_sfa_tmem_cols
|
||||
+ offset,
|
||||
dtype=self.sf_dtype,
|
||||
)
|
||||
@@ -1350,10 +1388,17 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
]
|
||||
|
||||
# Get accumulator stage index
|
||||
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
|
||||
|
||||
# Set tensor memory buffer for current tile
|
||||
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
|
||||
tTR_tAcc = tTR_tAcc_base[
|
||||
(None, None, None, None, None, acc_consumer_state.index)
|
||||
(None, None, None, None, None, acc_stage_index)
|
||||
]
|
||||
|
||||
#
|
||||
@@ -1370,12 +1415,27 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
|
||||
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
|
||||
for subtile_idx in cutlass.range(subtile_cnt):
|
||||
real_subtile_idx = subtile_idx
|
||||
if cutlass.const_expr(self.overlapping_accum):
|
||||
if reverse_subtile:
|
||||
real_subtile_idx = self.cta_tile_shape_mnk[1] // self.epi_tile_n - 1 - subtile_idx
|
||||
#
|
||||
# Load accumulator from tensor memory buffer to register
|
||||
#
|
||||
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
|
||||
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
|
||||
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
|
||||
|
||||
#
|
||||
# Async arrive accumulator buffer empty ealier when overlapping_accum is enabled
|
||||
#
|
||||
if cutlass.const_expr(self.overlapping_accum):
|
||||
if subtile_idx == self.iter_acc_early_release_in_epilogue:
|
||||
# Fence for TMEM load
|
||||
cute.arch.fence_view_async_tmem_load()
|
||||
with cute.arch.elect_one():
|
||||
acc_pipeline.consumer_release(acc_consumer_state)
|
||||
acc_consumer_state.advance()
|
||||
|
||||
#
|
||||
# Convert to C type
|
||||
#
|
||||
@@ -1386,7 +1446,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# Store C to shared memory
|
||||
#
|
||||
c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
|
||||
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
|
||||
cute.copy(
|
||||
tiled_copy_r2s,
|
||||
tRS_rC,
|
||||
@@ -1406,7 +1466,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
cute.copy(
|
||||
tma_atom_c,
|
||||
bSG_sC[(None, c_buffer)],
|
||||
bSG_gC[(None, subtile_idx)],
|
||||
bSG_gC[(None, real_subtile_idx)],
|
||||
)
|
||||
# Fence and barrier to make sure shared memory store is visible to TMA store
|
||||
c_pipeline.producer_commit()
|
||||
@@ -1416,9 +1476,10 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# Async arrive accumulator buffer empty
|
||||
#
|
||||
with cute.arch.elect_one():
|
||||
acc_pipeline.consumer_release(acc_consumer_state)
|
||||
acc_consumer_state.advance()
|
||||
if cutlass.const_expr(not self.overlapping_accum):
|
||||
with cute.arch.elect_one():
|
||||
acc_pipeline.consumer_release(acc_consumer_state)
|
||||
acc_consumer_state.advance()
|
||||
|
||||
#
|
||||
# Advance to next tile
|
||||
@@ -2286,6 +2347,7 @@ def run(
|
||||
c_tensor,
|
||||
max_active_clusters,
|
||||
current_stream,
|
||||
options=f"--opt-level 2",
|
||||
)
|
||||
|
||||
# Compute reference result
|
||||
|
||||
Reference in New Issue
Block a user