v4.3.1 update. (#2817)
This commit is contained in:
@@ -1977,33 +1977,21 @@ class BlockwiseGemmKernel:
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
if cutlass.const_expr(self.mma_tiler[0] == 64):
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
|
||||
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
|
||||
tAcc_final_epi = cute.flat_divide(
|
||||
|
||||
@@ -2010,33 +2010,21 @@ class BlockwiseContiguousGroupedGemmKernel:
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
if cutlass.const_expr(self.mma_tiler[0] == 64):
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
|
||||
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
|
||||
tAcc_final_epi = cute.flat_divide(
|
||||
|
||||
@@ -2010,33 +2010,21 @@ class BlockwiseMaskedGroupedGemmKernel:
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_load_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.Ld16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
if cutlass.const_expr(self.mma_tiler[0] == 64):
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(8)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
elif cutlass.const_expr(self.mma_tiler[0] == 128):
|
||||
else:
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St32x32bOp(tcgen05.copy.Repetition(32)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
else:
|
||||
# default: 16dp
|
||||
tmem_store_atom = cute.make_copy_atom(
|
||||
tcgen05.copy.St16x256bOp(tcgen05.copy.Repetition(1)),
|
||||
self.acc_dtype,
|
||||
)
|
||||
|
||||
tAcc_epi = cute.flat_divide(tAcc[((None, None), 0, 0, None)], epi_tile)
|
||||
tAcc_final_epi = cute.flat_divide(
|
||||
|
||||
Reference in New Issue
Block a user