CUTLASS 3.3.0 (#1167)

* Release 3.3.0

Adds support for mixed precision GEMMs On Hopper and Ampere
Adds support for < 16B aligned GEMMs on Hopper
Enhancements to EVT
Enhancements to Python interface
Enhancements to Sub-byte type handling in CuTe
Several other bug-fixes and performance improvements.

* minor doc update
This commit is contained in:
Pradeep Ramani
2023-11-02 11:09:05 -04:00
committed by GitHub
parent 922fb5108b
commit c008b4aea8
263 changed files with 16214 additions and 5008 deletions
@@ -294,6 +294,8 @@ struct DefaultMmaTensorOp<
Policy, PartitionsK, AccumulatorsInRowMajor>;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace warp
} // namespace gemm
} // namespace cutlass
@@ -132,10 +132,10 @@ public:
static int const kNumElementsInWarpFragment = NumElementsInWarpFragment;
static int const kNumElementsInMmaFragment = NumElementsInMmaFragment;
static Operand const kOperand = Operand::kA;
using WarpFragment = Array<ElementLoad, kNumElementsInWarpFragment>;
using MmaFragment = Array<ElementLoad, kNumElementsInMmaFragment>;
static uint32_t const kSelectBytesEvenThread = 0x5410;
static uint32_t const kSelectBytesOddThread = 0x7632;
@@ -168,7 +168,7 @@ public:
uint32_t const* src_ptr = reinterpret_cast<uint32_t const *>(&mma_frag_src_ptr[n]);
uint32_t *dst_ptr = reinterpret_cast<uint32_t *>(&mma_frag_dst_ptr[n]);
// Shuffle data within the warp, pull from other threads within the warp
uint32_t tmp0 = __shfl_up_sync(0xFFFFFFFF, src_ptr[0], delta_up_);
uint32_t tmp1 = __shfl_down_sync(0xFFFFFFFF, src_ptr[0], delta_down_);
@@ -218,7 +218,7 @@ public:
using WarpFragment = Array<ElementLoad, kNumElementsInWarpFragment>;
using MmaFragment = Array<ElementLoad, kNumElementsInMmaFragment>;
static uint32_t const kSelectBytesEvenThread = 0x5410;
static uint32_t const kSelectBytesOddThread = 0x7632;
@@ -260,7 +260,7 @@ public:
// Reorder the data within the 32-bit word (4x8b) required for mma.sync
dst_ptr[0] = __byte_perm(tmp0, tmp1, byte_selector_);
}
return result;
}
@@ -279,7 +279,7 @@ template <
///
typename Enable = void>
struct FragmentConverter {
using ElementDst = ElementDst_;
using ElementSrc = ElementSrc_;
@@ -522,17 +522,6 @@ public:
void transform(TransformedFragmentA &dst_A, TransformedFragmentB &dst_B,
FragmentA const &A, FragmentB const &B) const {
// Shuffle data within warp to obtain the mma.sync operand layout
detail::FragmentShuffler<MmaElementA, ElementA, MmaIterations::kRow,
FragmentA::kElements, MmaOperandA::kElements, Operand::kA> shuffler_A;
FragmentA tmp_A;
tmp_A = shuffler_A(A);
// Convert the A operand to the Mma Instruction operand type
detail::FragmentConverter<MmaElementA, ElementA, FragmentA::kElements> convert_A;
dst_A = convert_A(tmp_A);
// Shuffle data within warp to obtain the mma.sync operand layout
detail::FragmentShuffler<MmaElementB, ElementB, MmaIterations::kColumn,
FragmentB::kElements, MmaOperandB::kElements, Operand::kB> shuffler_B;
@@ -542,6 +531,27 @@ public:
// Convert the B operand to the Mma Instruction operand type
detail::FragmentConverter<MmaElementB, ElementB, FragmentB::kElements> convert_B;
dst_B = convert_B(tmp_B);
FragmentA tmp_A;
Array<ElementA, FragmentA::kElements / 2> *
ptr_tmp_A = reinterpret_cast<Array<ElementA,
FragmentA::kElements / 2> *>(&tmp_A);
Array<MmaElementA, FragmentA::kElements / 2> *
ptr_dst_A = reinterpret_cast<Array<MmaElementA,
FragmentA::kElements / 2> *>(&dst_A);
// Shuffle data within warp to obtain the mma.sync operand layout
detail::FragmentShuffler<MmaElementA, ElementA, MmaIterations::kRow,
FragmentA::kElements, MmaOperandA::kElements, Operand::kA> shuffler_A;
// Convert the A operand to the Mma Instruction operand type
detail::FragmentConverter<MmaElementA, ElementA, FragmentA::kElements / 2> convert_A;
tmp_A = shuffler_A(A);
ptr_dst_A[0] = convert_A(ptr_tmp_A[0]);
ptr_dst_A[1] = convert_A(ptr_tmp_A[1]);
}
};
@@ -551,4 +561,4 @@ public:
} // namespace gemm
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -158,6 +158,7 @@ public:
/// Max ID2
static int const kMaxID2 = Policy::Operator::kMaxID2;
static int const kVerticalVisit = false;
/// Data type of meta E that is moved at the same time
using ElementE =
typename cutlass::platform::conditional<kMaxID2 == 1, uint32_t,
@@ -251,8 +252,6 @@ public:
using MmaOperandC = typename Policy::Operator::FragmentC;
using MmaOperandE = typename Policy::Operator::FragmentE;
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
D = C;
MmaOperandA const *ptr_A = reinterpret_cast<MmaOperandA const *>(&A);
@@ -260,6 +259,36 @@ public:
MmaOperandC *ptr_D = reinterpret_cast<MmaOperandC *>(&D);
MmaOperandE const *ptr_E = reinterpret_cast<MmaOperandE const *>(&E);
if (kVerticalVisit) {
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < MmaIterations::kColumn; ++n) {
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < MmaIterations::kRow; ++m) {
int m_serpentine = ((n % 2) ? (MmaIterations::kRow - 1 - m) : m);
int id2 = m_serpentine % kMaxID2;
if (AccumulatorsInRowMajor) { // matrix B is reordered
mma(
ptr_D[n + m_serpentine * MmaIterations::kColumn],
ptr_A[m_serpentine],
ptr_B[n],
ptr_D[n + m_serpentine * MmaIterations::kColumn],
ptr_E[(m_serpentine / kMaxID2)],
id2);
} else {
mma(
ptr_D[m_serpentine + n * MmaIterations::kRow],
ptr_A[m_serpentine],
ptr_B[n],
ptr_D[m_serpentine + n * MmaIterations::kRow],
ptr_E[(m_serpentine / kMaxID2)],
id2);
}
}
}
} else {
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < MmaIterations::kRow; ++m) {
@@ -288,9 +317,7 @@ public:
}
}
}
#else
assert(0);
#endif
}
}
/// Transform the mma operands to the required types
@@ -298,7 +325,6 @@ public:
void transform(TransformedFragmentA &dst_A, TransformedFragmentB &dst_B,
FragmentA const &A, FragmentB const &B) const {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
//
// Define conversions from source type to instruction type
//
@@ -308,25 +334,42 @@ public:
FloatRoundStyle const kRoundB =
PreferredRoundingMode<typename ArchMmaOperator::ElementB,
ElementB>::kRound;
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
FragmentA::kElements / 2, kRoundA>
convert_A;
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
FragmentB::kElements, kRoundB>
convert_B;
Array<ElementA, FragmentA::kElements / 2> const *ptr_A =
reinterpret_cast<Array<ElementA, FragmentA::kElements / 2> const *>(&A);
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements / 2> *
ptr_dst_A = reinterpret_cast<Array<typename ArchMmaOperator::ElementA,
FragmentA::kElements / 2> *>(&dst_A);
dst_B = convert_B(B);
ptr_dst_A[0] = convert_A(ptr_A[0]);
ptr_dst_A[1] = convert_A(ptr_A[1]);
#else
assert(0);
#endif
if (kVerticalVisit) {
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
FragmentA::kElements, kRoundA>
convert_A;
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
FragmentB::kElements / 2, kRoundB>
convert_B;
Array<ElementB, FragmentB::kElements / 2> const *ptr_B =
reinterpret_cast<Array<ElementB, FragmentB::kElements / 2> const *>(&B);
Array<typename ArchMmaOperator::ElementB, FragmentB::kElements / 2> *
ptr_dst_B = reinterpret_cast<Array<typename ArchMmaOperator::ElementB,
FragmentB::kElements / 2> *>(&dst_B);
dst_A = convert_A(A);
ptr_dst_B[0] = convert_B(ptr_B[0]);
ptr_dst_B[1] = convert_B(ptr_B[1]);
} else {
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
FragmentA::kElements / 2, kRoundA>
convert_A;
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
FragmentB::kElements, kRoundB>
convert_B;
Array<ElementA, FragmentA::kElements / 2> const *ptr_A =
reinterpret_cast<Array<ElementA, FragmentA::kElements / 2> const *>(&A);
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements / 2> *
ptr_dst_A = reinterpret_cast<Array<typename ArchMmaOperator::ElementA,
FragmentA::kElements / 2> *>(&dst_A);
dst_B = convert_B(B);
ptr_dst_A[0] = convert_A(ptr_A[0]);
ptr_dst_A[1] = convert_A(ptr_A[1]);
}
}
};
+13 -29
View File
@@ -217,6 +217,12 @@ public:
/// Number of partitions along K dimension
static int const kPartitionsK = PartitionsK_;
#if defined(__CUDA_ARCH__) && ((__CUDA_ARCH__ < 800) || (__CUDA_ARCH__ == 890))
static int const kVerticalVisit = true;
#else
static int const kVerticalVisit = false;
#endif
public:
/// Iterates over the A operand in memory
@@ -293,16 +299,8 @@ public:
MmaOperandB const *ptr_B = reinterpret_cast<MmaOperandB const *>(&B);
MmaOperandC *ptr_D = reinterpret_cast<MmaOperandC *>(&D);
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
// Serpentine visitation order maximizing reuse of Rb
// The visitation order is like
// _
// | | | |
// | | | |
// |_| |_|
//
// Down Up Down Up
if (kVerticalVisit) {
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < MmaIterations::kColumn; ++n) {
@@ -326,16 +324,7 @@ public:
}
}
}
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
// Serpentine visitation order maximizing reuse of Ra
// The visitation order is like
// _________
// _________|
// |_________
// __________|
//
// Right Left Right Left
} else {
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < MmaIterations::kRow; ++m) {
@@ -358,9 +347,7 @@ public:
}
}
}
#else
assert(0);
#endif
}
}
/// Transform the mma operands to the required types
@@ -377,7 +364,7 @@ public:
FloatRoundStyle const kRoundB =
PreferredRoundingMode<typename ArchMmaOperator::ElementB,
ElementB>::kRound;
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
if (kVerticalVisit) {
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
FragmentA::kElements, kRoundA>
convert_A;
@@ -394,8 +381,7 @@ public:
ptr_dst_B[0] = convert_B(ptr_B[0]);
ptr_dst_B[1] = convert_B(ptr_B[1]);
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
} else {
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
FragmentA::kElements / 2, kRoundA>
convert_A;
@@ -412,9 +398,7 @@ public:
ptr_dst_A[0] = convert_A(ptr_A[0]);
ptr_dst_A[1] = convert_A(ptr_A[1]);
#else
assert(0);
#endif
}
}
};