Support for Mixed Input TensorOp (#1084)
* Passing warp-level mixed input F16*(S8/U8) tests * passing device-level mixed input F16*(S8/U8) tests * add to profiler - I8 (111 TFLOPs), U (123 TFLOPs) * fast numeric conversions (I8 = 132 TFLOPs, U8 = 148 TFLOPs) * Speedup reference compilation (REVERT THIS COMMIT) * wider_add.u32_packed_sub.f16x2 (I8 = 132TFLOP/s, U8 = 170 TFLOP/s) * Improve s8->f16 cvt and support bf16*u8 @158 TFLOPs * BF16 * S8 (142 TFLOPs) * Handle mixed-input upcast on OperandA (Support [S8|U8]*[F16|BF16] * rename OpMultiplyAddMixedInput to OpMultiplyAddMixedInputUpcast * Add device-level test and profiler support for upcast on operand A * Move shfl before the cvt and reduce #shfls by 1/2 * fix smem_usage calculation for mixed_input types * uncomment the stuff (getting ready for merge) * profiler changes and mixed-input reference * mixed input reference are in a new file * use platform instead of std * comments and typo only * Use CreateGemmOperator and delete CreateMixedInputGemmOperator * copyright for new files * rebase follow-up
This commit is contained in:
@@ -2340,7 +2340,8 @@ struct NumericArrayConverter<uint4b_t, int, N, Round> {
|
||||
/// Conversion operator for Array. See the comments before
|
||||
/// FastLinearCombinationClamp.
|
||||
template <typename T, typename S, int N,
|
||||
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest>
|
||||
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest,
|
||||
typename Enable = void>
|
||||
struct FastNumericArrayConverter {
|
||||
using result_type = Array<T, N>;
|
||||
using source_type = Array<S, N>;
|
||||
@@ -2441,6 +2442,225 @@ struct FastNumericArrayConverter<int8_t, float, N, Round> {
|
||||
result_type operator()(source_type const &s) const { return convert(s); }
|
||||
};
|
||||
|
||||
|
||||
/// Partial specialization for Array<cutlass::half_t, 4> <= Array<int8_t, 4>
|
||||
template <FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<cutlass::half_t, int8_t, 4, Round> {
|
||||
using result_type = Array<cutlass::half_t, 4>;
|
||||
using source_type = Array<int8_t, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
result_type result;
|
||||
|
||||
#if 0 // Scalar conversion (Please keep this code for reference for vectorized version below)
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
int16_t tmp = source[i] + 26112 /* 0x6600 */;
|
||||
result[i] = reinterpret_cast<cutlass::half_t const &>(tmp) - 1536.0_hf;
|
||||
}
|
||||
#endif
|
||||
|
||||
// Vectorized s8->f16 conversion using packed instructions
|
||||
uint32_t const* source_ptr = reinterpret_cast<uint32_t const*>(&source);
|
||||
uint32_t* result_ptr = reinterpret_cast<uint32_t*>(&result);
|
||||
|
||||
// Pack s8x2 (s8[1], s8[0]) -> s16x2 (sext.s8[1], sext.s8[0])
|
||||
// (See https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-prmt)
|
||||
// The inline ptx below uses `msb=0` and `msb=1` from the above link to sign extend the sign-bit in 0, 1, 2, 3 bytes of s8x4
|
||||
// into result_ptr[0] and result_ptr[1]'s 08-15 and 24-31 bits, respectively.
|
||||
// Note that `__byte_perm(source_ptr[0], source_ptr[0], 0x9180);` won't achieve the same and doesn't sign extend the sign-bit.
|
||||
// Thus, we use inline ptx `prmt.b32` instruction for the desired sign extend from `s8x2` to `s16x2`.
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(result_ptr[0]) : "r"(source_ptr[0]), "n"(0x9180));
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(result_ptr[1]) : "r"(source_ptr[0]), "n"(0xB3A2));
|
||||
|
||||
// In the absense of add.s16x2 instruction, use bit-wise operation to execute signed addition with magic numbers to achieve
|
||||
// the same result as add.s16x2 instruction.
|
||||
// (See https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#logic-and-shift-instructions-lop3)
|
||||
// For a logical operation F(a, b, c) the value of kImmLut can be computed by applying the same operation to
|
||||
// three predefined constant values as follows:
|
||||
// ta = 0xF0;
|
||||
// tb = 0xCC;
|
||||
// tc = 0xAA;
|
||||
// kImmLut = F(ta, tb, tc);
|
||||
// If we want F = ((a & b) ^ c) then set kImmLut = (0xF0 & 0xCC) ^ 0xAA
|
||||
static constexpr uint32_t kImmLut = (0xF0 & 0xCC) ^ 0xAA;
|
||||
|
||||
// The bit-wise operation executed below is `result_ptr[0] = (result_ptr[0] & 0x03FF03FF) ^ 0x66006600;`
|
||||
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n" :
|
||||
"=r"(result_ptr[0]) : "r"(result_ptr[0]), "n"(0x03FF03FF), "n"(0x66006600), "n"(kImmLut));
|
||||
// The bit-wise operation executed below is `result_ptr[1] = (result_ptr[1] & 0x03FF03FF) ^ 0x66006600;`
|
||||
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n" :
|
||||
"=r"(result_ptr[1]) : "r"(result_ptr[1]), "n"(0x03FF03FF), "n"(0x66006600), "n"(kImmLut));
|
||||
|
||||
// Packed sub.f16x2 with magic number to obtain final converted result
|
||||
asm volatile("sub.f16x2 %0, %1, %2;\n" : "=r"(result_ptr[0]) : "r"(result_ptr[0]), "r"(0x66006600));
|
||||
asm volatile("sub.f16x2 %0, %1, %2;\n" : "=r"(result_ptr[1]) : "r"(result_ptr[1]), "r"(0x66006600));
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/// Partial specialization for Array<cutlass::half_t, 4> <= Array<uint8_t, 4>
|
||||
template <FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<cutlass::half_t, uint8_t, 4, Round> {
|
||||
using result_type = Array<cutlass::half_t, 4>;
|
||||
using source_type = Array<uint8_t, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
result_type result;
|
||||
|
||||
uint32_t const* source_ptr = reinterpret_cast<uint32_t const*>(&source);
|
||||
uint32_t* result_ptr = reinterpret_cast<uint32_t*>(&result);
|
||||
|
||||
result_ptr[0] = __byte_perm(source_ptr[0], 0x0, 0x4140);
|
||||
result_ptr[1] = __byte_perm(source_ptr[0], 0x0, 0x4342);
|
||||
|
||||
asm volatile("add.u32 %0, %1, %2;\n" : "=r"(result_ptr[0]) : "r"(result_ptr[0]), "r"(0x66006600));
|
||||
asm volatile("add.u32 %0, %1, %2;\n" : "=r"(result_ptr[1]) : "r"(result_ptr[1]), "r"(0x66006600));
|
||||
|
||||
asm volatile("sub.f16x2 %0, %1, %2;\n" : "=r"(result_ptr[0]) : "r"(result_ptr[0]), "r"(0x66006600));
|
||||
asm volatile("sub.f16x2 %0, %1, %2;\n" : "=r"(result_ptr[1]) : "r"(result_ptr[1]), "r"(0x66006600));
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<cutlass::bfloat16_t, 4> <= Array<uint8_t, 4>
|
||||
template <FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<cutlass::bfloat16_t, uint8_t, 4, Round> {
|
||||
using result_type = Array<cutlass::bfloat16_t, 4>;
|
||||
using source_type = Array<uint8_t, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
result_type result;
|
||||
Array<float, 4> tmp;
|
||||
|
||||
uint32_t const* source_ptr = reinterpret_cast<uint32_t const*>(&source);
|
||||
uint32_t* tmp_ptr = reinterpret_cast<uint32_t*>(&tmp);
|
||||
|
||||
// __byte_perm simulates the add.u32 0x4B000000 to every u8 element of u8x4 source and stores
|
||||
// the result in tmp (without introducing extra cvt.u32.u8 instruction)
|
||||
tmp_ptr[0] = __byte_perm(source_ptr[0], 0x4B000000, 0x7650);
|
||||
tmp_ptr[1] = __byte_perm(source_ptr[0], 0x4B000000, 0x7651);
|
||||
tmp_ptr[2] = __byte_perm(source_ptr[0], 0x4B000000, 0x7652);
|
||||
tmp_ptr[3] = __byte_perm(source_ptr[0], 0x4B000000, 0x7653);
|
||||
|
||||
// Subtract the magic number 0x4B000000 from tmp in floating-point arithmetic to obtain final result
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
tmp[i] = reinterpret_cast<float const &>(tmp_ptr[i]) - 8388608.f;
|
||||
}
|
||||
|
||||
// on 3456x4096x8192 runs at 158 TFLOP/s
|
||||
// Convert f32x2 to bf16x2 using `cvt.rn.b16x2.f32` instruction
|
||||
NumericArrayConverter<cutlass::bfloat16_t, float, 4, Round> convert_f32_to_bf16;
|
||||
result = convert_f32_to_bf16(tmp);
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<cutlass::bfloat16_t, 4> <= Array<int8_t, 4>
|
||||
template <FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<cutlass::bfloat16_t, int8_t, 4, Round> {
|
||||
using result_type = Array<cutlass::bfloat16_t, 4>;
|
||||
using source_type = Array<int8_t, 4>;
|
||||
|
||||
using intermediate_float_type = Array<float, 4>;
|
||||
using intermediate_int32_type = Array<int32_t, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
result_type result;
|
||||
intermediate_float_type tmp;
|
||||
|
||||
uint32_t const* source_ptr = reinterpret_cast<uint32_t const*>(&source);
|
||||
uint32_t* tmp_ptr = reinterpret_cast<uint32_t*>(&tmp);
|
||||
|
||||
// s8x4 (s[3], s[2], s8[1], s8[0]) -> s16x4 (sext.s8[3], sext.s8[2], sext.s8[1], sext.s8[0])
|
||||
// (See https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-prmt)
|
||||
// The inline ptx below uses `msb=0` and `msb=1` from the above link to sext the sign-bit in 0, 1, 2, 3 bytes of s8x4
|
||||
// sext without unpacking each s8 out of s8x4 into a separate register a.ka. without using shifts (SHFL).
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(tmp_ptr[0]) : "r"(source_ptr[0]), "n"(0x8880));
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(tmp_ptr[1]) : "r"(source_ptr[0]), "n"(0x9991));
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(tmp_ptr[2]) : "r"(source_ptr[0]), "n"(0xAAA2));
|
||||
asm volatile("prmt.b32 %0,%1,%1,%2;\n" : "=r"(tmp_ptr[3]) : "r"(source_ptr[0]), "n"(0xBBB3));
|
||||
|
||||
// Convert s32x4 to f32x4 using fast numeric array converter
|
||||
FastNumericArrayConverter<float, int32_t, 4, Round> convert_s32_to_f32_;
|
||||
tmp = convert_s32_to_f32_(reinterpret_cast<intermediate_int32_type const &>(tmp[0]));
|
||||
|
||||
// Convert f32x2 to bf16x2 using `cvt.rn.b16x2.f32` instruction
|
||||
NumericArrayConverter<cutlass::bfloat16_t, float, 4, Round> convert_f32_to_bf16_;
|
||||
result = convert_f32_to_bf16_(tmp);
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for FastNumericArrayConverter to vectorize over 4 elements.
|
||||
/// source `S` as 8b integers (S8 or U8) -> destination `T` as 16b floating-point (F16 or BF16)
|
||||
template <typename T, typename S, int N, FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<T, S, N, Round,
|
||||
typename platform::enable_if<(platform::is_same<T, half_t>::value || platform::is_same<T, bfloat16_t>::value) &&
|
||||
(platform::is_same<S, int8_t>::value || platform::is_same<S, uint8_t>::value)>::type> {
|
||||
static_assert(!(N % 4), "N must be multiple of 4.");
|
||||
|
||||
using result_type = Array<T, N>;
|
||||
using source_type = Array<S, N>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
FastNumericArrayConverter<T, S, 4, Round> convert_vector_;
|
||||
result_type result;
|
||||
|
||||
Array<T, 4> *result_ptr =
|
||||
reinterpret_cast<Array<T, 4> *>(&result);
|
||||
Array<S, 4> const *source_ptr =
|
||||
reinterpret_cast<Array<S, 4> const *>(&source);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N / 4; ++i) {
|
||||
result_ptr[i] = convert_vector_(source_ptr[i]);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const { return convert(s); }
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines preferred rounding mode for a pair of types
|
||||
|
||||
Reference in New Issue
Block a user