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:
@@ -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]);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user