[Bug Fix]Set NumSplitsM to 1 when TileShapeM < 128 in sm90 fp8 blockwise scaling CollectiveMma (#2965)
* Fix NumSplitsM when TileShapeM < 128. * Use cute::conditional_t to replace std::conditional_t.
This commit is contained in:
@@ -212,7 +212,8 @@ struct CollectiveMma<
|
||||
|
||||
static_assert(cute::is_same_v<ElementAccumulator, ElementBlockScale>,
|
||||
"ElementAccumulator and ElementBlockScale should be same datatype");
|
||||
using NumSplitsM = cute::C<get<0>(TileShape_{}) / 128>;
|
||||
// For TileShapeM < 128, NumSplitsM should be 1
|
||||
using NumSplitsM = cute::conditional_t<get<0>(TileShape_{}) < _128{}, _1, cute::C<get<0>(TileShape_{}) / 128>>;
|
||||
static_assert(NumSplitsM{} == 1 || NumSplitsM{} == 2);
|
||||
|
||||
struct SharedStorage {
|
||||
|
||||
Reference in New Issue
Block a user