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:
@@ -1303,11 +1303,195 @@ struct NumericArrayConverter<uint8_t, int, N, Round> {
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<float, 2> <= Array<float_e4m3_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float, float_e4m3_t, 2, Round> {
|
||||
using result_element = float;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
using result_type = Array<result_element, 2>;
|
||||
using source_type = Array<source_element, 2>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
#if defined(CUDA_PTX_FP8_CVT_ENABLED)
|
||||
uint32_t out_fp16;
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e4m3x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(out_fp16): "h"(src_packed));
|
||||
|
||||
float2 res0 = __half22float2(reinterpret_cast<__half2 &>(out_fp16));
|
||||
|
||||
result_type out;
|
||||
out[0] = res0.x;
|
||||
out[1] = res0.y;
|
||||
return out;
|
||||
#else
|
||||
result_type result;
|
||||
NumericConverter<result_element, source_element, Round> converter;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
result[i] = converter(source[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<float_e4m3_t, 2> <= Array<float, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, float, 2, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = float;
|
||||
|
||||
using result_type = Array<result_element, 2>;
|
||||
using source_type = Array<source_element, 2>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
#if defined(CUDA_PTX_FP8_CVT_ENABLED)
|
||||
uint16_t out;
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 %0, %2, %1;\n" \
|
||||
"}" \
|
||||
: "=h"(out) : "f"(source[0]), "f"(source[1]));
|
||||
|
||||
return reinterpret_cast<result_type const &>(out);
|
||||
#else
|
||||
result_type result;
|
||||
NumericConverter<result_element, source_element, Round> converter;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
result[i] = converter(source[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<float, 2> <= Array<float_e5m2_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float, float_e5m2_t, 2, Round> {
|
||||
using result_element = float;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
using result_type = Array<result_element, 2>;
|
||||
using source_type = Array<source_element, 2>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
#if defined(CUDA_PTX_FP8_CVT_ENABLED)
|
||||
uint32_t out_fp16;
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e5m2x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(out_fp16): "h"(src_packed));
|
||||
|
||||
float2 res0 = __half22float2(reinterpret_cast<__half2 &>(out_fp16));
|
||||
|
||||
result_type out;
|
||||
out[0] = res0.x;
|
||||
out[1] = res0.y;
|
||||
return out;
|
||||
#else
|
||||
result_type result;
|
||||
NumericConverter<result_element, source_element, Round> converter;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
result[i] = converter(source[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
namespace detail {
|
||||
|
||||
/// Special converters that can be used with 4 8-bit elements packed in a register.
|
||||
/// Common use is for fast FP8 converters.
|
||||
template <
|
||||
typename T,
|
||||
typename S,
|
||||
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest,
|
||||
typename Transform = cutlass::transform::thread::UnaryTransform::Identity
|
||||
>
|
||||
struct NumericArrayConverterPacked4Element {
|
||||
using result_type = Array<T, 4>;
|
||||
using source_type = Array<S, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
static_assert(platform::is_same<Transform, cutlass::transform::thread::UnaryTransform::Identity>::value ||
|
||||
platform::is_same<Transform, cutlass::transform::thread::UnaryTransform::Conjugate>::value,
|
||||
"Unary Operator not supported.");
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & s) {
|
||||
|
||||
result_type result;
|
||||
NumericConverter<T, S, Round> convert_;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
if (platform::is_same<Transform, cutlass::transform::thread::UnaryTransform::Identity>::value) {
|
||||
result[i] = convert_(s[i]);
|
||||
}
|
||||
else { // conjugate
|
||||
result[i] = conj(convert_(s[i]));
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<float, 4> <= Array<float_e4m3_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float, float_e4m3_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float, float_e4m3_t, Round> {
|
||||
using result_element = float;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
@@ -1362,7 +1546,7 @@ struct NumericArrayConverter<float, float_e4m3_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, float, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e4m3_t, float, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = float;
|
||||
|
||||
@@ -1406,11 +1590,17 @@ struct NumericArrayConverter<float_e4m3_t, float, 4, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for Array<float, 4> <=> Array<float_e5m2_t, 4>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<float, 4> <= Array<float_e5m2_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float, float_e5m2_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float, float_e5m2_t, Round> {
|
||||
using result_element = float;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
@@ -1465,7 +1655,7 @@ struct NumericArrayConverter<float, float_e5m2_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, float, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e5m2_t, float, Round> {
|
||||
using result_element = float_e5m2_t;
|
||||
using source_element = float;
|
||||
|
||||
@@ -1519,7 +1709,7 @@ struct NumericArrayConverter<float_e5m2_t, float, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<half_t, float_e4m3_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<half_t, float_e4m3_t, Round> {
|
||||
using result_element = half_t;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
@@ -1564,7 +1754,7 @@ struct NumericArrayConverter<half_t, float_e4m3_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, half_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e4m3_t, half_t, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = half_t;
|
||||
|
||||
@@ -1609,11 +1799,17 @@ struct NumericArrayConverter<float_e4m3_t, half_t, 4, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for Array<half_t, 4> <=> Array<float_e5m2_t, 4>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<half_t, 4> <= Array<float_e5m2_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<half_t, float_e5m2_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<half_t, float_e5m2_t, Round> {
|
||||
using result_element = half_t;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
@@ -1658,7 +1854,7 @@ struct NumericArrayConverter<half_t, float_e5m2_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, half_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e5m2_t, half_t, Round> {
|
||||
using result_element = float_e5m2_t;
|
||||
using source_element = half_t;
|
||||
|
||||
@@ -1713,7 +1909,7 @@ struct NumericArrayConverter<float_e5m2_t, half_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<bfloat16_t, float_e4m3_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<bfloat16_t, float_e4m3_t, Round> {
|
||||
using result_element = bfloat16_t;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
@@ -1726,7 +1922,7 @@ struct NumericArrayConverter<bfloat16_t, float_e4m3_t, 4, Round> {
|
||||
|
||||
#if defined(CUDA_PTX_FP8_CVT_ENABLED)
|
||||
// Convert f8 to float
|
||||
NumericArrayConverter<float, source_element, 4, Round> src2float;
|
||||
NumericArrayConverterPacked4Element<float, source_element, Round> src2float;
|
||||
Array<float, 4> tmp_floats = src2float(source);
|
||||
|
||||
// Convert float to bf16
|
||||
@@ -1761,7 +1957,7 @@ struct NumericArrayConverter<bfloat16_t, float_e4m3_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, bfloat16_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e4m3_t, bfloat16_t, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = bfloat16_t;
|
||||
|
||||
@@ -1782,7 +1978,7 @@ struct NumericArrayConverter<float_e4m3_t, bfloat16_t, 4, Round> {
|
||||
packed_tmp[1] = src2float(packed_source[1]);
|
||||
|
||||
// Convert float to f8
|
||||
NumericArrayConverter<result_element, float, 4, Round> float2result;
|
||||
NumericArrayConverterPacked4Element<result_element, float, Round> float2result;
|
||||
return float2result(tmp);
|
||||
#else
|
||||
result_type result;
|
||||
@@ -1803,11 +1999,17 @@ struct NumericArrayConverter<float_e4m3_t, bfloat16_t, 4, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for Array<bfloat16_t, 4> <=> Array<float_e5m2_t, 4>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<bfloat16_t, 4> <= Array<float_e5m2_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<bfloat16_t, float_e5m2_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<bfloat16_t, float_e5m2_t, Round> {
|
||||
using result_element = bfloat16_t;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
@@ -1820,7 +2022,7 @@ struct NumericArrayConverter<bfloat16_t, float_e5m2_t, 4, Round> {
|
||||
|
||||
#if defined(CUDA_PTX_FP8_CVT_ENABLED)
|
||||
// Convert f8 to float
|
||||
NumericArrayConverter<float, source_element, 4, Round> src2float;
|
||||
NumericArrayConverterPacked4Element<float, source_element, Round> src2float;
|
||||
Array<float, 4> tmp_floats = src2float(source);
|
||||
|
||||
// Convert float to bf16
|
||||
@@ -1855,7 +2057,7 @@ struct NumericArrayConverter<bfloat16_t, float_e5m2_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, bfloat16_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e5m2_t, bfloat16_t, Round> {
|
||||
using result_element = float_e5m2_t;
|
||||
using source_element = bfloat16_t;
|
||||
|
||||
@@ -1876,7 +2078,7 @@ struct NumericArrayConverter<float_e5m2_t, bfloat16_t, 4, Round> {
|
||||
packed_tmp[1] = src2float(packed_source[1]);
|
||||
|
||||
// Convert float to f8
|
||||
NumericArrayConverter<result_element, float, 4, Round> float2result;
|
||||
NumericArrayConverterPacked4Element<result_element, float, Round> float2result;
|
||||
return float2result(tmp);
|
||||
#else
|
||||
result_type result;
|
||||
@@ -1907,7 +2109,7 @@ struct NumericArrayConverter<float_e5m2_t, bfloat16_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, float_e5m2_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e4m3_t, float_e5m2_t, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
@@ -1938,7 +2140,7 @@ struct NumericArrayConverter<float_e4m3_t, float_e5m2_t, 4, Round> {
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, float_e4m3_t, 4, Round> {
|
||||
struct NumericArrayConverterPacked4Element<float_e5m2_t, float_e4m3_t, Round> {
|
||||
using result_element = float_e5m2_t;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
@@ -1965,63 +2167,7 @@ struct NumericArrayConverter<float_e5m2_t, float_e4m3_t, 4, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specializations for:
|
||||
// Array<float_e4m3_t, 4> <=> Array<float_e4m3_t, 4>
|
||||
// Array<float_e5m2_t, 4> <=> Array<float_e5m2_t, 4>
|
||||
//
|
||||
// These are needed to avoid multiple-matching-template compilation errors (e.g., when
|
||||
// compiling float_e4m3_t <=> float_e4m3_t, which among T <= float_e4m3_t and float_e4m3_t <= T
|
||||
// should be used?)
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<float_e4m3_t, 4> <= Array<float_e4m3_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, float_e4m3_t, 4, Round> {
|
||||
using result_element = float_e4m3_t;
|
||||
using source_element = float_e4m3_t;
|
||||
|
||||
using result_type = Array<result_element, 4>;
|
||||
using source_type = Array<source_element, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
return source;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<float_e5m2_t, 4> <= Array<float_e5m2_t, 4>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, float_e5m2_t, 4, Round> {
|
||||
using result_element = float_e5m2_t;
|
||||
using source_element = float_e5m2_t;
|
||||
|
||||
using result_type = Array<result_element, 4>;
|
||||
using source_type = Array<source_element, 4>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const &source) {
|
||||
return source;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
@@ -2058,7 +2204,7 @@ public:
|
||||
packed_result_type* packed_result = reinterpret_cast<packed_result_type*>(&result);
|
||||
const packed_source_type* packed_source = reinterpret_cast<const packed_source_type*>(&source);
|
||||
|
||||
NumericArrayConverter<result_element, source_element, 4, Round> packed_converter;
|
||||
detail::NumericArrayConverterPacked4Element<result_element, source_element, Round> packed_converter;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N / 4; ++i) {
|
||||
@@ -2150,8 +2296,11 @@ template <
|
||||
struct NumericArrayConverter<float_e5m2_t, float_e5m2_t, N, Round> :
|
||||
public PackedNumericArrayConverter<float_e5m2_t, float_e5m2_t, N, Round> {};
|
||||
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
/// Partial specialization for Array<int8_t> <= Array<float>
|
||||
/// Conversion is performed with saturation regardless of setting of
|
||||
/// the `Round` template parameter.
|
||||
@@ -2360,7 +2509,9 @@ struct FastNumericArrayConverter {
|
||||
|
||||
/// Partial specialization for Array<float> <= Array<int>
|
||||
template <typename T, int N, FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<float, T, N, Round> {
|
||||
struct FastNumericArrayConverter<float, T, N, Round,
|
||||
typename platform::enable_if<platform::numeric_limits<T>::is_integer>
|
||||
> {
|
||||
using result_type = Array<float, N>;
|
||||
using source_type = Array<T, N>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
@@ -2442,7 +2593,6 @@ 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> {
|
||||
@@ -2454,7 +2604,7 @@ struct FastNumericArrayConverter<cutlass::half_t, int8_t, 4, 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) {
|
||||
@@ -2471,8 +2621,8 @@ struct FastNumericArrayConverter<cutlass::half_t, int8_t, 4, Round> {
|
||||
// (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`.
|
||||
// Note that `__byte_perm(source_ptr[0], source_ptr[0], 0x9180);` won't acheive 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));
|
||||
|
||||
@@ -2508,7 +2658,6 @@ struct FastNumericArrayConverter<cutlass::half_t, int8_t, 4, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/// Partial specialization for Array<cutlass::half_t, 4> <= Array<uint8_t, 4>
|
||||
template <FloatRoundStyle Round>
|
||||
struct FastNumericArrayConverter<cutlass::half_t, uint8_t, 4, Round> {
|
||||
@@ -2519,7 +2668,7 @@ struct FastNumericArrayConverter<cutlass::half_t, uint8_t, 4, 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);
|
||||
|
||||
@@ -2632,7 +2781,7 @@ struct FastNumericArrayConverter<cutlass::bfloat16_t, int8_t, 4, Round> {
|
||||
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> {
|
||||
(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>;
|
||||
@@ -2658,7 +2807,7 @@ struct FastNumericArrayConverter<T, S, N, Round,
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const { return convert(s); }
|
||||
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user