CUTLASS 2.2 (#96)
Adds support for NVIDIA Ampere Architecture features. CUDA 11 Toolkit recommended.
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -45,7 +45,9 @@ enum class FloatRoundStyle {
|
||||
round_toward_zero, ///< round toward zero
|
||||
round_to_nearest, ///< round to nearest even
|
||||
round_toward_infinity, ///< round toward infinity
|
||||
round_toward_neg_infinity ///< round toward negative infinity
|
||||
round_toward_neg_infinity, ///< round toward negative infinity
|
||||
round_half_ulp_truncate, ///< add 0.5ulp to integer representation then round toward zero
|
||||
round_half_ulp_trunc_dntz ///< like round_half_ulp_truncate, except denorms are rounded *toward* zero
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -240,6 +242,232 @@ struct NumericConverter<half_t, float, FloatRoundStyle::round_toward_zero> {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for float <=> bfloat16_t
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for float <= bfloat16_t
|
||||
template <FloatRoundStyle Round>
|
||||
struct NumericConverter<float, bfloat16_t, Round> {
|
||||
|
||||
using result_type = float;
|
||||
using source_type = bfloat16_t;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
return static_cast<float>(s);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<bfloat16_t, float, FloatRoundStyle::round_to_nearest> {
|
||||
using result_type = bfloat16_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_to_nearest;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
return static_cast<bfloat16_t>(s);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<bfloat16_t, float, FloatRoundStyle::round_half_ulp_truncate> {
|
||||
using result_type = bfloat16_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_half_ulp_truncate;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
uint32_t x32 = reinterpret_cast<uint32_t const &>(s);
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
if (::isfinite(s)) {
|
||||
x32 += 0x8000;
|
||||
}
|
||||
#else
|
||||
if (std::isfinite(s)) {
|
||||
x32 += 0x8000;
|
||||
}
|
||||
#endif
|
||||
|
||||
uint16_t x16 = uint16_t((x32 >> 16) & 0xffff);
|
||||
return bfloat16_t::bitcast(x16);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<bfloat16_t, float, FloatRoundStyle::round_toward_zero> {
|
||||
using result_type = bfloat16_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_toward_zero;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
uint32_t x32 = reinterpret_cast<uint32_t const &>(s);
|
||||
uint16_t x16 = uint16_t(x32 >> 16);
|
||||
|
||||
return bfloat16_t::bitcast(x16);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for float <=> tfloat32_t
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for float <= tfloat32_t
|
||||
template <FloatRoundStyle Round>
|
||||
struct NumericConverter<float, tfloat32_t, Round> {
|
||||
|
||||
using result_type = float;
|
||||
using source_type = tfloat32_t;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
return static_cast<float>(s);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<tfloat32_t, float, FloatRoundStyle::round_to_nearest> {
|
||||
using result_type = tfloat32_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_to_nearest;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
unsigned storage = reinterpret_cast<unsigned const &>(s);
|
||||
|
||||
if ((storage & 0x7f800000) != 0x7f800000) {
|
||||
|
||||
bool mantissa_bit = ((storage & (1 << 13)) != 0);
|
||||
bool round_bit = ((storage & (1 << 12)) != 0);
|
||||
bool sticky_bit = ((storage & ((1 << 12) - 1)) != 0);
|
||||
|
||||
if ((round_bit && sticky_bit) || (round_bit && mantissa_bit)) {
|
||||
storage += uint32_t(1 << 13);
|
||||
}
|
||||
|
||||
// Note, the following is intentionally commented out. TF32
|
||||
// does not define the low order bits, so they may be left in
|
||||
// an undefined state.
|
||||
//
|
||||
// By not truncating these bit explicitly, we avoid an extra logical
|
||||
// operation.
|
||||
//
|
||||
// TF32 may be implicitly converted to float by performing this
|
||||
// operation as needed.
|
||||
//
|
||||
// storage = (storage & ~0x1fff);
|
||||
}
|
||||
else if (storage & ~0xff800000) {
|
||||
storage = 0x7fffffff;
|
||||
}
|
||||
|
||||
return tfloat32_t::bitcast(storage);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<tfloat32_t, float, FloatRoundStyle::round_half_ulp_truncate> {
|
||||
using result_type = tfloat32_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_half_ulp_truncate;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
return tfloat32_t::round_half_ulp_truncate(s);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// This rounding operation is similar to half_ulp_truncate except it rounds denorms toward zero.
|
||||
/// It avoids predicated code, though it requires a temporary register.
|
||||
template <>
|
||||
struct NumericConverter<tfloat32_t, float, FloatRoundStyle::round_half_ulp_trunc_dntz> {
|
||||
using result_type = tfloat32_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_half_ulp_trunc_dntz;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
unsigned y = reinterpret_cast<unsigned const &>(s);
|
||||
y = y & 0xff800000;
|
||||
float d = reinterpret_cast<float const &>(y);
|
||||
float z = d / float(1 << 11) + s;
|
||||
|
||||
return reinterpret_cast<result_type const &>(z);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct NumericConverter<tfloat32_t, float, FloatRoundStyle::round_toward_zero> {
|
||||
using result_type = tfloat32_t;
|
||||
using source_type = float;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_toward_zero;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
uint32_t x = reinterpret_cast<uint32_t const &>(s);
|
||||
return tfloat32_t::bitcast(x & 0xffffe000);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Conversion and Clamp operator for Integers
|
||||
@@ -518,6 +746,77 @@ struct NumericArrayConverter<float, half_t, N, Round> {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<bfloat16_t, 2> <= Array<float, 2>, round to nearest
|
||||
template <>
|
||||
struct NumericArrayConverter<bfloat16_t, float, 2, FloatRoundStyle::round_to_nearest> {
|
||||
|
||||
using result_type = Array<bfloat16_t, 2>;
|
||||
using source_type = Array<float, 2>;
|
||||
static FloatRoundStyle const round_style = FloatRoundStyle::round_to_nearest;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
unsigned d;
|
||||
|
||||
asm("cvt.rn.bf16x2.f32 %0, %1, %2;\n" : "=r"(d) : "f"(source[1]), "f"(source[0]) );
|
||||
|
||||
return reinterpret_cast<result_type const &>(d);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<half> <= Array<float>
|
||||
template <
|
||||
int N,
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<bfloat16_t, float, N, Round> {
|
||||
|
||||
using result_type = Array<bfloat16_t, N>;
|
||||
using source_type = Array<float, N>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
NumericArrayConverter<bfloat16_t, float, 2, Round> convert_vector_;
|
||||
NumericConverter<bfloat16_t, float, Round> convert_element_;
|
||||
|
||||
result_type result;
|
||||
|
||||
Array<bfloat16_t, 2> *result_ptr = reinterpret_cast<Array<bfloat16_t, 2> *>(&result);
|
||||
Array<float, 2> const *source_ptr = reinterpret_cast<Array<float, 2> const *>(&source);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N / 2; ++i) {
|
||||
result_ptr[i] = convert_vector_(source_ptr[i]);
|
||||
}
|
||||
|
||||
if (N % 2) {
|
||||
result[N - 1] = convert_element_(source[N - 1]);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
#endif // if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Conditional guards to enable partial specialization for packed integers
|
||||
@@ -843,6 +1142,12 @@ struct PreferredRoundingMode {
|
||||
static FloatRoundStyle const kRound = FloatRoundStyle::round_to_nearest;
|
||||
};
|
||||
|
||||
/// Defines preferred rounding mode for a pair of types
|
||||
template <>
|
||||
struct PreferredRoundingMode<tfloat32_t, float> {
|
||||
static FloatRoundStyle const kRound = FloatRoundStyle::round_half_ulp_truncate;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass
|
||||
|
||||
Reference in New Issue
Block a user