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:
@@ -27,7 +27,7 @@
|
||||
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
|
||||
import argparse
|
||||
from typing import Optional, Type, Tuple, Union
|
||||
from typing import Type, Tuple, Union
|
||||
|
||||
import cuda.bindings.driver as cuda
|
||||
import torch
|
||||
@@ -85,7 +85,7 @@ Input arguments to this example is shown below:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
python examples/blackwell/dense_blockscaled_gemm_persistent.py \
|
||||
python examples/blackwell/dense_blockscaled_gemm_persistent.py \
|
||||
--ab_dtype Float4E2M1FN --sf_dtype Float8E8M0FNU --sf_vec_size 16 \
|
||||
--c_dtype Float16 \
|
||||
--mma_tiler_mn 256,128 --cluster_shape_mn 2,1 \
|
||||
@@ -95,7 +95,7 @@ To collect performance with NCU profiler:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
ncu python examples/blackwell/dense_blockscaled_gemm_persistent.py \
|
||||
ncu python examples/blackwell/dense_blockscaled_gemm_persistent.py \
|
||||
--ab_dtype Float4E2M1FN --sf_dtype Float8E8M0FNU --sf_vec_size 16 \
|
||||
--c_dtype Float16 \
|
||||
--mma_tiler_mn 256,128 --cluster_shape_mn 2,1 \
|
||||
@@ -108,7 +108,7 @@ Constraints:
|
||||
see detailed valid dtype combinations in below Sm100BlockScaledPersistentDenseGemmKernel class documentation
|
||||
* A/B tensor must have the same data type, mixed data type is not supported (e.g., mxf8 x mxf4)
|
||||
* Mma tiler M must be 128 or 256(use_2cta_instrs)
|
||||
* Mma tiler N must be 128 or 256
|
||||
* Mma tiler N must be 64/128/192/256
|
||||
* Cluster shape M/N must be positive and power of 2, total cluster size <= 16
|
||||
* Cluster shape M must be multiple of 2 if Mma tiler M is 256(use_2cta_instrs)
|
||||
* The contiguous dimension of A/B/C tensors must be at least 16 bytes aligned,
|
||||
@@ -144,7 +144,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
- Float8E4M3FN/Float8E5M2
|
||||
:note: Constraints:
|
||||
- MMA tiler M must be 128 or 256 (use_2cta_instrs)
|
||||
- MMA tiler N must be 128/256
|
||||
- MMA tiler N must be 64/128/192/256
|
||||
- Cluster shape M must be multiple of 2 if Mma tiler M is 256
|
||||
- Cluster shape M/N must be positive and power of 2, total cluster size <= 16
|
||||
- Also, Cluster shape M/N must be <= 4 for scale factor multicasts due to limited size of scale factors
|
||||
@@ -209,9 +209,18 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
(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_bar_id = 0
|
||||
self.epilog_sync_bar_id = 1
|
||||
self.tmem_ptr_sync_bar_id = 2
|
||||
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)),
|
||||
)
|
||||
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
|
||||
@@ -228,21 +237,17 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
- Computing epilogue subtile
|
||||
- Setting up A/B/SFA/SFB/C stage counts in shared memory
|
||||
- Computing A/B/SFA/SFB/C shared memory layout
|
||||
- Computing tensor memory allocation columns
|
||||
"""
|
||||
# Compute mma instruction shapes
|
||||
mma_inst_bits_k = 256
|
||||
# (MMA_Tile_Shape_M, MMA_Tile_Shape_N, MMA_Inst_Shape_K)
|
||||
self.mma_inst_shape_mnk = (
|
||||
self.mma_inst_shape_mn = (
|
||||
self.mma_tiler[0],
|
||||
self.mma_tiler[1],
|
||||
mma_inst_bits_k // self.a_dtype.width,
|
||||
)
|
||||
# (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K)
|
||||
self.mma_inst_shape_mnk_sfb = (
|
||||
self.mma_inst_shape_mnk[0] // (2 if self.use_2cta_instrs else 1),
|
||||
cute.round_up(self.mma_inst_shape_mnk[1], 128),
|
||||
self.mma_inst_shape_mnk[2],
|
||||
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(
|
||||
@@ -252,7 +257,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.sf_dtype,
|
||||
self.sf_vec_size,
|
||||
self.cta_group,
|
||||
self.mma_inst_shape_mnk[:2],
|
||||
self.mma_inst_shape_mn,
|
||||
)
|
||||
|
||||
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
|
||||
@@ -262,20 +267,21 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.sf_dtype,
|
||||
self.sf_vec_size,
|
||||
cute.nvgpu.tcgen05.CtaGroup.ONE,
|
||||
self.mma_inst_shape_mnk_sfb[:2],
|
||||
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_mnk[0],
|
||||
self.mma_inst_shape_mnk[1],
|
||||
self.mma_inst_shape_mnk[2] * mma_inst_tile_k,
|
||||
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_mnk_sfb[0],
|
||||
self.mma_inst_shape_mnk_sfb[1],
|
||||
self.mma_inst_shape_mnk_sfb[2] * mma_inst_tile_k,
|
||||
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),
|
||||
@@ -314,9 +320,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
tiled_mma,
|
||||
self.mma_tiler,
|
||||
self.a_dtype,
|
||||
self.a_major_mode,
|
||||
self.b_dtype,
|
||||
self.b_major_mode,
|
||||
self.epi_tile,
|
||||
self.c_dtype,
|
||||
self.c_layout,
|
||||
@@ -431,7 +435,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.sf_dtype,
|
||||
self.sf_vec_size,
|
||||
self.cta_group,
|
||||
self.mma_inst_shape_mnk[:2],
|
||||
self.mma_inst_shape_mn,
|
||||
)
|
||||
|
||||
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
|
||||
@@ -441,7 +445,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
self.sf_dtype,
|
||||
self.sf_vec_size,
|
||||
cute.nvgpu.tcgen05.CtaGroup.ONE,
|
||||
self.mma_inst_shape_mnk_sfb[:2],
|
||||
self.mma_inst_shape_mn_sfb,
|
||||
)
|
||||
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
|
||||
|
||||
@@ -507,6 +511,31 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
internal_type=cutlass.Int16,
|
||||
)
|
||||
|
||||
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
|
||||
x = tma_tensor_sfb.stride[0][1]
|
||||
y = cute.ceil_div(tma_tensor_sfb.shape[0][1], 4)
|
||||
|
||||
new_shape = (
|
||||
(
|
||||
tma_tensor_sfb.shape[0][0],
|
||||
((2, 2), y)
|
||||
),
|
||||
tma_tensor_sfb.shape[1],
|
||||
tma_tensor_sfb.shape[2]
|
||||
)
|
||||
# Use right multiplication for ScaledBasis (3 * x instead of x * 3)
|
||||
x_times_3 = 3 * x
|
||||
new_stride = (
|
||||
(
|
||||
tma_tensor_sfb.stride[0][0],
|
||||
((x, x), x_times_3)
|
||||
),
|
||||
tma_tensor_sfb.stride[1],
|
||||
tma_tensor_sfb.stride[2]
|
||||
)
|
||||
tma_tensor_sfb_new_layout = cute.make_layout(new_shape, stride=new_stride)
|
||||
tma_tensor_sfb = cute.make_tensor(tma_tensor_sfb.iterator, tma_tensor_sfb_new_layout)
|
||||
|
||||
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
|
||||
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
|
||||
sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
|
||||
@@ -628,7 +657,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
mSFA_mkl: cute.Tensor,
|
||||
tma_atom_sfb: cute.CopyAtom,
|
||||
mSFB_nkl: cute.Tensor,
|
||||
tma_atom_c: Optional[cute.CopyAtom],
|
||||
tma_atom_c: cute.CopyAtom,
|
||||
mC_mnl: cute.Tensor,
|
||||
cluster_layout_vmnk: cute.Layout,
|
||||
cluster_layout_sfb_vmnk: cute.Layout,
|
||||
@@ -636,7 +665,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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, None],
|
||||
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
|
||||
epi_tile: cute.Tile,
|
||||
tile_sched_params: utils.PersistentTileSchedulerParams,
|
||||
epilogue_op: cutlass.Constexpr,
|
||||
@@ -684,9 +713,6 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
smem = utils.SmemAllocator()
|
||||
storage = smem.allocate(self.shared_storage)
|
||||
|
||||
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
|
||||
tmem_holding_buf = storage.tmem_holding_buf
|
||||
|
||||
# Initialize mainloop ab_pipeline (barrier) and states
|
||||
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
|
||||
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
|
||||
@@ -719,14 +745,13 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
# 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
|
||||
)
|
||||
cute.arch.mbarrier_init_fence()
|
||||
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,
|
||||
)
|
||||
|
||||
# Cluster arrive after barrier init
|
||||
if cute.size(self.cluster_shape_mn) > 1:
|
||||
@@ -790,7 +815,9 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
# (bN, bK, RestN, RestK, RestL)
|
||||
gSFB_nkl = cute.local_tile(
|
||||
mSFB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
|
||||
mSFB_nkl,
|
||||
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
|
||||
(None, None, None),
|
||||
)
|
||||
# (bM, bN, RestM, RestN, RestL)
|
||||
gC_mnl = cute.local_tile(
|
||||
@@ -894,9 +921,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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()
|
||||
|
||||
#
|
||||
# Specialized TMA load warp
|
||||
@@ -915,7 +940,6 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
while work_tile.is_valid_tile:
|
||||
|
||||
# Get tile coord from tile scheduler
|
||||
cur_tile_coord = work_tile.tile_idx
|
||||
mma_tile_coord_mnl = (
|
||||
@@ -940,9 +964,13 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
tAgSFA_slice = tAgSFA[
|
||||
(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
|
||||
# ((atom_v, rest_v), RestK)
|
||||
tBgSFB_slice = tBgSFB[
|
||||
(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
|
||||
(None, slice_n, None, mma_tile_coord_mnl[2])
|
||||
]
|
||||
|
||||
# Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt
|
||||
@@ -1017,21 +1045,13 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# Bar sync for retrieve tensor memory 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/SFA/SFB tensor
|
||||
#
|
||||
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
|
||||
# Make accumulator tmem tensor
|
||||
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
|
||||
self.acc_dtype,
|
||||
alignment=16,
|
||||
ptr_to_buffer_holding_addr=tmem_holding_buf,
|
||||
)
|
||||
# (MMA, MMA_M, MMA_N, STAGE)
|
||||
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
|
||||
|
||||
@@ -1067,12 +1087,16 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# Partition for S2T copy of SFA/SFB
|
||||
#
|
||||
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)
|
||||
)
|
||||
(
|
||||
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)
|
||||
|
||||
#
|
||||
# Persistent tile scheduling loop
|
||||
@@ -1116,6 +1140,30 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
|
||||
+ 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
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtAcc_base)
|
||||
+ tcgen05.find_tmem_tensor_col_offset(tCtSFA)
|
||||
+ offset,
|
||||
dtype=self.sf_dtype,
|
||||
)
|
||||
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
|
||||
|
||||
#
|
||||
# Reset the ACCUMULATE field for each tile
|
||||
#
|
||||
@@ -1170,7 +1218,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
tiled_mma.set(
|
||||
tcgen05.Field.SFB,
|
||||
tCtSFB[sf_kblock_coord].iterator,
|
||||
tCtSFB_mma[sf_kblock_coord].iterator,
|
||||
)
|
||||
|
||||
cute.gemm(
|
||||
@@ -1220,30 +1268,17 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# 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
|
||||
#
|
||||
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
|
||||
self.acc_dtype,
|
||||
alignment=16,
|
||||
ptr_to_buffer_holding_addr=tmem_holding_buf,
|
||||
)
|
||||
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
|
||||
# (MMA, MMA_M, MMA_N, STAGE)
|
||||
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
|
||||
|
||||
@@ -1251,20 +1286,24 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
# Partition for epilogue
|
||||
#
|
||||
epi_tidx = tidx
|
||||
tiled_copy_t2r, tTR_tAcc_base, tTR_rAcc = (
|
||||
self.epilog_tmem_copy_and_partition(
|
||||
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
|
||||
)
|
||||
(
|
||||
tiled_copy_t2r,
|
||||
tTR_tAcc_base,
|
||||
tTR_rAcc,
|
||||
) = self.epilog_tmem_copy_and_partition(
|
||||
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
|
||||
)
|
||||
tma_atom_c, bSG_sC, bSG_gC_partitioned = (
|
||||
self.epilog_gmem_copy_and_partition(
|
||||
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
|
||||
)
|
||||
(
|
||||
tma_atom_c,
|
||||
bSG_sC,
|
||||
bSG_gC_partitioned,
|
||||
) = self.epilog_gmem_copy_and_partition(
|
||||
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
|
||||
)
|
||||
|
||||
#
|
||||
@@ -1283,7 +1322,6 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
c_producer_group = pipeline.CooperativeGroup(
|
||||
pipeline.Agent.Thread,
|
||||
32 * len(self.epilog_warp_id),
|
||||
32 * len(self.epilog_warp_id),
|
||||
)
|
||||
c_pipeline = pipeline.PipelineTmaStore.create(
|
||||
num_stages=self.num_c_stage,
|
||||
@@ -1291,7 +1329,6 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
while work_tile.is_valid_tile:
|
||||
|
||||
# Get tile coord from tile scheduler
|
||||
cur_tile_coord = work_tile.tile_idx
|
||||
mma_tile_coord_mnl = (
|
||||
@@ -1360,11 +1397,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
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()
|
||||
|
||||
#
|
||||
# TMA store C to global memory
|
||||
@@ -1378,10 +1411,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
# Fence and barrier to make sure shared memory store is visible to TMA store
|
||||
c_pipeline.producer_commit()
|
||||
c_pipeline.producer_acquire()
|
||||
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
|
||||
@@ -1399,21 +1429,9 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
#
|
||||
# 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(
|
||||
acc_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(acc_tmem_ptr)
|
||||
#
|
||||
# Wait for C store complete
|
||||
#
|
||||
@@ -1520,7 +1538,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
# (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
|
||||
@@ -1614,9 +1632,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
tiled_mma: cute.TiledMma,
|
||||
mma_tiler_mnk: Tuple[int, int, int],
|
||||
a_dtype: Type[cutlass.Numeric],
|
||||
a_major_mode: tcgen05.OperandMajorMode,
|
||||
b_dtype: Type[cutlass.Numeric],
|
||||
b_major_mode: tcgen05.OperandMajorMode,
|
||||
epi_tile: cute.Tile,
|
||||
c_dtype: Type[cutlass.Numeric],
|
||||
c_layout: utils.LayoutEnum,
|
||||
@@ -1633,12 +1649,8 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
:type mma_tiler_mnk: tuple[int, int, int]
|
||||
:param a_dtype: Data type of operand A.
|
||||
:type a_dtype: type[cutlass.Numeric]
|
||||
:param a_major_mode: Major mode of operand A.
|
||||
:type a_major_mode: tcgen05.OperandMajorMode
|
||||
:param b_dtype: Data type of operand B.
|
||||
:type b_dtype: type[cutlass.Numeric]
|
||||
:param b_major_mode: Major mode of operand B.
|
||||
:type b_major_mode: tcgen05.OperandMajorMode
|
||||
:param epi_tile: The epilogue tile shape.
|
||||
:type epi_tile: cute.Tile
|
||||
:param c_dtype: Data type of operand C (output).
|
||||
@@ -1830,7 +1842,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
c_major: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if the dtypes and sf_vec_size are valid combinations
|
||||
Check if layouts and dtypes are valid combinations
|
||||
|
||||
:param ab_dtype: The data type of the A and B operands
|
||||
:type ab_dtype: Type[cutlass.Numeric]
|
||||
@@ -1870,9 +1882,9 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
|
||||
"""
|
||||
is_valid = True
|
||||
# Skip invalid mma tile shape
|
||||
if not mma_tiler_mn[0] in [128, 256]:
|
||||
if mma_tiler_mn[0] not in [128, 256]:
|
||||
is_valid = False
|
||||
if not mma_tiler_mn[1] in [128, 256]:
|
||||
if mma_tiler_mn[1] not in [64, 128, 192, 256]:
|
||||
is_valid = False
|
||||
# Skip illegal cluster shape
|
||||
if cluster_shape_mn[0] % (2 if mma_tiler_mn[0] == 256 else 1) != 0:
|
||||
@@ -2088,7 +2100,7 @@ def run(
|
||||
:return: Execution time of the GEMM kernel
|
||||
:rtype: float
|
||||
"""
|
||||
print(f"Running Sm100 Persistent Dense BlockScaled GEMM test with:")
|
||||
print("Running Sm100 Persistent Dense BlockScaled GEMM test with:")
|
||||
print(f"mnkl: {mnkl}")
|
||||
print(f"AB dtype: {ab_dtype}, SF dtype: {sf_dtype}, SF Vec size: {sf_vec_size}")
|
||||
print(f"C dtype: {c_dtype}")
|
||||
@@ -2143,21 +2155,21 @@ def run(
|
||||
c_ref, c_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
)
|
||||
|
||||
# Mark tensor to be byte aligned
|
||||
# Mark tensor with element divisibility for 16B alignment
|
||||
a_tensor.mark_compact_shape_dynamic(
|
||||
mode=1 if a_major == "k" else 0,
|
||||
stride_order=(2, 0, 1) if a_major == "k" else (2, 1, 0),
|
||||
divisibility=2 if ab_dtype == cutlass.Float4E2M1FN else 1,
|
||||
divisibility=32 if ab_dtype == cutlass.Float4E2M1FN else 16,
|
||||
)
|
||||
b_tensor.mark_compact_shape_dynamic(
|
||||
mode=1 if b_major == "k" else 0,
|
||||
stride_order=(2, 0, 1) if b_major == "k" else (2, 1, 0),
|
||||
divisibility=2 if ab_dtype == cutlass.Float4E2M1FN else 1,
|
||||
divisibility=32 if ab_dtype == cutlass.Float4E2M1FN else 16,
|
||||
)
|
||||
c_tensor.mark_compact_shape_dynamic(
|
||||
mode=1 if c_major == "n" else 0,
|
||||
stride_order=(2, 0, 1) if c_major == "n" else (2, 1, 0),
|
||||
divisibility=2 if c_dtype == cutlass.Float4E2M1FN else 1,
|
||||
divisibility=32 if ab_dtype == cutlass.Float4E2M1FN else 16,
|
||||
)
|
||||
|
||||
# Create scale factor tensor SFA/SFB
|
||||
@@ -2374,6 +2386,7 @@ def run(
|
||||
|
||||
return exec_time # Return execution time in microseconds
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
def parse_comma_separated_ints(s: str) -> Tuple[int, ...]:
|
||||
|
||||
Reference in New Issue
Block a user