v4.3.1 update. (#2817)

This commit is contained in:
Junkai-Wu
2025-11-27 09:49:30 -05:00
committed by GitHub
parent 2052fd3885
commit 1de3a576cc
44 changed files with 3316 additions and 510 deletions
@@ -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(