v4.1 release
This commit is contained in:
@@ -35,6 +35,7 @@ import torch
|
||||
import cutlass
|
||||
import cutlass.cute as cute
|
||||
import cutlass.utils as utils
|
||||
import cutlass.pipeline as pipeline
|
||||
from cutlass.cute.nvgpu import cpasync, tcgen05
|
||||
import cutlass.torch as cutlass_torch
|
||||
import cutlass.utils.blackwell_helpers as sm100_utils
|
||||
@@ -211,7 +212,7 @@ class DenseGemmKernel:
|
||||
|
||||
self.occupancy = 1
|
||||
self.threads_per_cta = 128
|
||||
self.num_smem_capacity = sm100_utils.SMEM_CAPACITY["sm100"]
|
||||
self.smem_capacity = sm100_utils.SMEM_CAPACITY["sm100"]
|
||||
|
||||
def _setup_attributes(self):
|
||||
"""Set up configurations that are dependent on GEMM inputs
|
||||
@@ -283,7 +284,7 @@ class DenseGemmKernel:
|
||||
self.epi_tile,
|
||||
self.c_dtype,
|
||||
self.c_layout,
|
||||
self.num_smem_capacity,
|
||||
self.smem_capacity,
|
||||
self.occupancy,
|
||||
self.use_tma_store,
|
||||
)
|
||||
@@ -308,7 +309,7 @@ class DenseGemmKernel:
|
||||
self.epi_tile,
|
||||
self.num_c_stage,
|
||||
)
|
||||
if cutlass.const_expr(self.use_tma_store)
|
||||
if self.use_tma_store
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -372,9 +373,11 @@ class DenseGemmKernel:
|
||||
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
|
||||
|
||||
# Setup TMA load for A
|
||||
a_op = self._get_tma_atom_kind(atom_thr_size, self.is_a_mcast)
|
||||
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))
|
||||
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tma_tile_atom_A(
|
||||
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
|
||||
a_op,
|
||||
a,
|
||||
a_smem_layout,
|
||||
@@ -387,9 +390,11 @@ class DenseGemmKernel:
|
||||
)
|
||||
|
||||
# Setup TMA load for B
|
||||
b_op = self._get_tma_atom_kind(atom_thr_size, self.is_b_mcast)
|
||||
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))
|
||||
tma_atom_b, tma_tensor_b = cute.nvgpu.make_tma_tile_atom_B(
|
||||
tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
|
||||
b_op,
|
||||
b,
|
||||
b_smem_layout,
|
||||
@@ -413,7 +418,7 @@ class DenseGemmKernel:
|
||||
cute.make_identity_layout(c.shape), self.epi_tile
|
||||
)
|
||||
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
|
||||
tma_atom_c, tma_tensor_c = cpasync.make_tma_tile_atom(
|
||||
tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
|
||||
cpasync.CopyBulkTensorTileS2GOp(),
|
||||
c,
|
||||
epi_smem_layout,
|
||||
@@ -426,9 +431,7 @@ class DenseGemmKernel:
|
||||
self.buffer_align_bytes = 1024
|
||||
|
||||
c_smem_size = (
|
||||
cute.cosize(self.c_smem_layout_staged.outer)
|
||||
if cutlass.const_expr(self.use_tma_store)
|
||||
else 0
|
||||
cute.cosize(self.c_smem_layout_staged.outer) if self.use_tma_store else 0
|
||||
)
|
||||
|
||||
# Define shared storage for kernel
|
||||
@@ -472,7 +475,7 @@ class DenseGemmKernel:
|
||||
tma_atom_b,
|
||||
tma_tensor_b,
|
||||
tma_atom_c,
|
||||
tma_tensor_c if cutlass.const_expr(self.use_tma_store) else c,
|
||||
tma_tensor_c if self.use_tma_store else c,
|
||||
self.cluster_layout_vmnk,
|
||||
self.a_smem_layout_staged,
|
||||
self.b_smem_layout_staged,
|
||||
@@ -556,12 +559,12 @@ class DenseGemmKernel:
|
||||
tmem_holding_buf = storage.tmem_holding_buf
|
||||
|
||||
# Initialize mainloop ab_pipeline (barrier) and states
|
||||
ab_pipeline_producer_group = utils.CooperativeGroup(utils.Agent.Thread)
|
||||
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 = utils.CooperativeGroup(
|
||||
utils.Agent.Thread, num_tma_producer
|
||||
ab_pipeline_consumer_group = pipeline.CooperativeGroup(
|
||||
pipeline.Agent.Thread, num_tma_producer
|
||||
)
|
||||
ab_pipeline = utils.PipelineTmaUmma.create(
|
||||
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,
|
||||
@@ -569,30 +572,30 @@ class DenseGemmKernel:
|
||||
tx_count=self.num_tma_load_bytes,
|
||||
cta_layout_vmnk=cluster_layout_vmnk,
|
||||
)
|
||||
ab_producer_state = utils.make_pipeline_state(
|
||||
utils.PipelineUserType.Producer, self.num_ab_stage
|
||||
ab_producer_state = pipeline.make_pipeline_state(
|
||||
pipeline.PipelineUserType.Producer, self.num_ab_stage
|
||||
)
|
||||
ab_consumer_state = utils.make_pipeline_state(
|
||||
utils.PipelineUserType.Consumer, self.num_ab_stage
|
||||
ab_consumer_state = pipeline.make_pipeline_state(
|
||||
pipeline.PipelineUserType.Consumer, self.num_ab_stage
|
||||
)
|
||||
|
||||
# Initialize acc_pipeline (barrier) and states
|
||||
acc_pipeline_producer_group = utils.CooperativeGroup(utils.Agent.Thread)
|
||||
acc_pipeline_consumer_group = utils.CooperativeGroup(
|
||||
utils.Agent.Thread, self.threads_per_cta, self.threads_per_cta
|
||||
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
|
||||
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
|
||||
pipeline.Agent.Thread, self.threads_per_cta, self.threads_per_cta
|
||||
)
|
||||
acc_pipeline = utils.PipelineUmmaAsync.create(
|
||||
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,
|
||||
)
|
||||
acc_producer_state = utils.make_pipeline_state(
|
||||
utils.PipelineUserType.Producer, self.num_acc_stage
|
||||
acc_producer_state = pipeline.make_pipeline_state(
|
||||
pipeline.PipelineUserType.Producer, self.num_acc_stage
|
||||
)
|
||||
acc_consumer_state = utils.make_pipeline_state(
|
||||
utils.PipelineUserType.Consumer, self.num_acc_stage
|
||||
acc_consumer_state = pipeline.make_pipeline_state(
|
||||
pipeline.PipelineUserType.Consumer, self.num_acc_stage
|
||||
)
|
||||
|
||||
# Tensor memory dealloc barrier init
|
||||
@@ -600,7 +603,7 @@ class DenseGemmKernel:
|
||||
if warp_idx == 0:
|
||||
num_tmem_dealloc_threads = 32
|
||||
with cute.arch.elect_one():
|
||||
cute.arch.mbarrier_init_arrive_cnt(
|
||||
cute.arch.mbarrier_init(
|
||||
tmem_dealloc_mbar_ptr, num_tmem_dealloc_threads
|
||||
)
|
||||
cute.arch.mbarrier_init_fence()
|
||||
@@ -617,7 +620,7 @@ class DenseGemmKernel:
|
||||
storage.sC.get_tensor(
|
||||
c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
|
||||
)
|
||||
if cutlass.const_expr(self.use_tma_store)
|
||||
if self.use_tma_store
|
||||
else None
|
||||
)
|
||||
# (MMA, MMA_M, MMA_K, STAGE)
|
||||
@@ -634,7 +637,7 @@ class DenseGemmKernel:
|
||||
#
|
||||
a_full_mcast_mask = None
|
||||
b_full_mcast_mask = None
|
||||
if self.is_a_mcast or self.is_b_mcast or use_2cta_instrs:
|
||||
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
|
||||
)
|
||||
@@ -645,15 +648,15 @@ class DenseGemmKernel:
|
||||
#
|
||||
# Local_tile partition global tensors
|
||||
#
|
||||
# (bM, bK, loopM, loopK, loopL)
|
||||
# (bM, bK, RestM, RestK, RestL)
|
||||
gA_mkl = cute.local_tile(
|
||||
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
|
||||
)
|
||||
# (bN, bK, loopN, loopK, loopL)
|
||||
# (bN, bK, RestN, RestK, RestL)
|
||||
gB_nkl = cute.local_tile(
|
||||
mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
|
||||
)
|
||||
# (bM, bN, loopM, loopN, loopL)
|
||||
# (bM, bN, RestM, RestN, RestL)
|
||||
gC_mnl = cute.local_tile(
|
||||
mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
|
||||
)
|
||||
@@ -663,11 +666,11 @@ class DenseGemmKernel:
|
||||
# Partition global tensor for TiledMMA_A/B/C
|
||||
#
|
||||
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
|
||||
# (MMA, MMA_M, MMA_K, loopM, loopK, loopL)
|
||||
# (MMA, MMA_M, MMA_K, RestM, RestK, RestL)
|
||||
tCgA = thr_mma.partition_A(gA_mkl)
|
||||
# (MMA, MMA_N, MMA_K, loopN, loopK, loopL)
|
||||
# (MMA, MMA_N, MMA_K, RestN, RestK, RestL)
|
||||
tCgB = thr_mma.partition_B(gB_nkl)
|
||||
# (MMA, MMA_M, MMA_N, loopM, loopN, loopL)
|
||||
# (MMA, MMA_M, MMA_N, RestM, RestN, RestL)
|
||||
tCgC = thr_mma.partition_C(gC_mnl)
|
||||
|
||||
#
|
||||
@@ -678,7 +681,7 @@ class DenseGemmKernel:
|
||||
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
|
||||
)
|
||||
# ((atom_v, rest_v), STAGE)
|
||||
# ((atom_v, rest_v), loopM, loopK, loopL)
|
||||
# ((atom_v, rest_v), RestM, RestK, RestL)
|
||||
tAsA, tAgA = cpasync.tma_partition(
|
||||
tma_atom_a,
|
||||
block_in_cluster_coord_vmnk[2],
|
||||
@@ -691,7 +694,7 @@ class DenseGemmKernel:
|
||||
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
|
||||
)
|
||||
# ((atom_v, rest_v), STAGE)
|
||||
# ((atom_v, rest_v), loopN, loopK, loopL)
|
||||
# ((atom_v, rest_v), RestN, RestK, RestL)
|
||||
tBsB, tBgB = cpasync.tma_partition(
|
||||
tma_atom_b,
|
||||
block_in_cluster_coord_vmnk[1],
|
||||
@@ -771,9 +774,9 @@ class DenseGemmKernel:
|
||||
#
|
||||
# Slice to per mma tile index
|
||||
#
|
||||
# ((atom_v, rest_v), loopK)
|
||||
# ((atom_v, rest_v), RestK)
|
||||
tAgA = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
|
||||
# ((atom_v, rest_v), loopK)
|
||||
# ((atom_v, rest_v), RestK)
|
||||
tBgB = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
|
||||
if cutlass.const_expr(self.use_tma_store):
|
||||
# ((ATOM_V, REST_V), EPI_M, EPI_N)
|
||||
@@ -797,7 +800,7 @@ class DenseGemmKernel:
|
||||
#
|
||||
# Prefetch TMA load A/B
|
||||
#
|
||||
for prefetch_idx in cutlass.range_dynamic(prefetch_k_block_cnt, unroll=1):
|
||||
for prefetch_idx in cutlass.range(prefetch_k_block_cnt, unroll=1):
|
||||
# Conditionally wait for AB buffer empty
|
||||
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
|
||||
|
||||
@@ -833,7 +836,7 @@ class DenseGemmKernel:
|
||||
#
|
||||
# MMA mainloop
|
||||
#
|
||||
for k_block in cutlass.range_dynamic(0, k_block_cnt, 1, unroll=1):
|
||||
for k_block in range(k_block_cnt):
|
||||
# Conditionally wait for AB buffer empty
|
||||
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
|
||||
|
||||
@@ -860,7 +863,7 @@ class DenseGemmKernel:
|
||||
|
||||
# tCtAcc += tCrA * tCrB
|
||||
num_kphases = cute.size(tCrA, mode=[2])
|
||||
for kphase_idx in range(num_kphases):
|
||||
for kphase_idx in cutlass.range(num_kphases, unroll_full=True):
|
||||
kphase_coord = (None, None, kphase_idx, ab_consumer_state.index)
|
||||
|
||||
cute.gemm(
|
||||
@@ -917,10 +920,10 @@ class DenseGemmKernel:
|
||||
c_pipeline = None
|
||||
if cutlass.const_expr(self.use_tma_store):
|
||||
# Initialize tma store c_pipeline
|
||||
c_producer_group = utils.CooperativeGroup(
|
||||
utils.Agent.Thread, self.threads_per_cta, self.threads_per_cta
|
||||
c_producer_group = pipeline.CooperativeGroup(
|
||||
pipeline.Agent.Thread, self.threads_per_cta, self.threads_per_cta
|
||||
)
|
||||
c_pipeline = utils.PipelineTmaStore.create(
|
||||
c_pipeline = pipeline.PipelineTmaStore.create(
|
||||
num_stages=self.num_c_stage,
|
||||
producer_group=c_producer_group,
|
||||
)
|
||||
@@ -929,7 +932,7 @@ class DenseGemmKernel:
|
||||
# Store accumulator to global memory in subtiles
|
||||
#
|
||||
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
|
||||
for subtile_idx in cutlass.range_dynamic(subtile_cnt):
|
||||
for subtile_idx in range(subtile_cnt):
|
||||
#
|
||||
# Load accumulator from tensor memory buffer to register
|
||||
#
|
||||
@@ -1007,7 +1010,7 @@ class DenseGemmKernel:
|
||||
#
|
||||
if warp_idx == 0:
|
||||
# Reverse prefetch_k_block_cnt times to next available buffer
|
||||
for i in cutlass.range_dynamic(prefetch_k_block_cnt):
|
||||
for i in range(prefetch_k_block_cnt):
|
||||
ab_producer_state.reverse()
|
||||
ab_pipeline.producer_tail(ab_producer_state)
|
||||
return
|
||||
@@ -1063,11 +1066,11 @@ class DenseGemmKernel:
|
||||
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
|
||||
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
|
||||
|
||||
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, loopM, loopN, loopL)
|
||||
# (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, loopM, loopN, loopL)
|
||||
# (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(
|
||||
@@ -1149,7 +1152,7 @@ class DenseGemmKernel:
|
||||
- tTR_gC: The partitioned global tensor C
|
||||
:rtype: Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]
|
||||
"""
|
||||
# (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, loopM, loopN, loopL)
|
||||
# (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
|
||||
)
|
||||
@@ -1158,7 +1161,7 @@ class DenseGemmKernel:
|
||||
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, loopM, loopN, loopL)
|
||||
# ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL)
|
||||
bSG_sC, bSG_gC = cpasync.tma_partition(
|
||||
tma_atom_c,
|
||||
0,
|
||||
@@ -1169,7 +1172,7 @@ class DenseGemmKernel:
|
||||
return tma_atom_c, bSG_sC, bSG_gC
|
||||
else:
|
||||
tiled_copy_t2r = atom
|
||||
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N, loopM, loopN, loopL)
|
||||
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL)
|
||||
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
|
||||
tTR_gC = thr_copy_t2r.partition_D(gC_epi)
|
||||
# (T2R, T2R_M, T2R_N)
|
||||
@@ -1188,7 +1191,7 @@ class DenseGemmKernel:
|
||||
epi_tile: cute.Tile,
|
||||
c_dtype: Type[cutlass.Numeric],
|
||||
c_layout: utils.LayoutEnum,
|
||||
num_smem_capacity: int,
|
||||
smem_capacity: int,
|
||||
occupancy: int,
|
||||
use_tma_store: bool,
|
||||
) -> Tuple[int, int, int]:
|
||||
@@ -1208,8 +1211,8 @@ class DenseGemmKernel:
|
||||
:type c_dtype: type[cutlass.Numeric]
|
||||
:param c_layout: Layout enum of operand C in global memory.
|
||||
:type c_layout: utils.LayoutEnum
|
||||
:param num_smem_capacity: Total available shared memory capacity in bytes.
|
||||
:type num_smem_capacity: int
|
||||
:param smem_capacity: Total available shared memory capacity in bytes.
|
||||
:type smem_capacity: int
|
||||
:param occupancy: Target number of CTAs per SM (occupancy).
|
||||
:type occupancy: int
|
||||
:param use_tma_store: Whether TMA store is enabled.
|
||||
@@ -1263,7 +1266,7 @@ class DenseGemmKernel:
|
||||
# Subtract reserved bytes and initial C stages bytes
|
||||
# Divide remaining by bytes needed per A/B stage
|
||||
num_ab_stage = (
|
||||
num_smem_capacity - (occupancy + 1) * (mbar_helpers_bytes + c_bytes)
|
||||
smem_capacity - (occupancy + 1) * (mbar_helpers_bytes + c_bytes)
|
||||
) // ab_bytes_per_stage
|
||||
|
||||
# Refine epilogue stages:
|
||||
@@ -1271,7 +1274,7 @@ class DenseGemmKernel:
|
||||
# Add remaining unused smem to epilogue
|
||||
if use_tma_store:
|
||||
num_c_stage += (
|
||||
num_smem_capacity
|
||||
smem_capacity
|
||||
- ab_bytes_per_stage * num_ab_stage
|
||||
- (occupancy + 1) * (mbar_helpers_bytes + c_bytes)
|
||||
) // ((occupancy + 1) * c_bytes_per_stage)
|
||||
@@ -1309,36 +1312,6 @@ class DenseGemmKernel:
|
||||
|
||||
return grid
|
||||
|
||||
@staticmethod
|
||||
def _get_tma_atom_kind(
|
||||
atom_sm_cnt: cutlass.Int32, mcast: cutlass.Boolean
|
||||
) -> Union[
|
||||
cpasync.CopyBulkTensorTileG2SMulticastOp, cpasync.CopyBulkTensorTileG2SOp
|
||||
]:
|
||||
"""
|
||||
Select the appropriate TMA copy atom based on the number of SMs and the multicast flag.
|
||||
|
||||
:param atom_sm_cnt: The number of SMs
|
||||
:type atom_sm_cnt: cutlass.Int32
|
||||
:param mcast: The multicast flag
|
||||
:type mcast: cutlass.Boolean
|
||||
|
||||
:return: The appropriate TMA copy atom kind
|
||||
:rtype: cpasync.CopyBulkTensorTileG2SMulticastOp or cpasync.CopyBulkTensorTileG2SOp
|
||||
|
||||
:raise ValueError: If the atom_sm_cnt is invalid
|
||||
"""
|
||||
if atom_sm_cnt == 2 and mcast:
|
||||
return cpasync.CopyBulkTensorTileG2SMulticastOp(tcgen05.CtaGroup.TWO)
|
||||
elif atom_sm_cnt == 2 and not mcast:
|
||||
return cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.TWO)
|
||||
elif atom_sm_cnt == 1 and mcast:
|
||||
return cpasync.CopyBulkTensorTileG2SMulticastOp(tcgen05.CtaGroup.ONE)
|
||||
elif atom_sm_cnt == 1 and not mcast:
|
||||
return cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE)
|
||||
|
||||
raise ValueError(f"Invalid atom_sm_cnt: {atom_sm_cnt} and {mcast}")
|
||||
|
||||
@staticmethod
|
||||
def _compute_num_tmem_alloc_cols(
|
||||
tiled_mma: cute.TiledMma, mma_tiler: Tuple[int, int, int]
|
||||
|
||||
Reference in New Issue
Block a user