CUTLASS 3.6.0 (#1850)
* v3.6 * update changelog * update readme * fix typo * fixing typos * hopper gemm with weight prefetch --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
parent
0837a2a00a
commit
cc3c29a81a
@@ -45,6 +45,7 @@
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/half.h"
|
||||
#include "cutlass/bfloat16.h"
|
||||
|
||||
namespace cutlass {
|
||||
|
||||
@@ -1602,6 +1603,423 @@ struct NumericArrayConverter<float, cutlass::float_e5m2_t, 2, Round> {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<float_e5m2_t, 2> <= Array<float, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, float, 2, Round> {
|
||||
using result_element = cutlass::float_e5m2_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.e5m2x2.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 specializations for Array<half, N> <=> Array<float_e4m3_t, N>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<half, 2> <= Array<float_e4m3_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<cutlass::half_t, cutlass::float_e4m3_t, 2, Round> {
|
||||
using result_element = cutlass::half_t;
|
||||
using source_element = cutlass::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)
|
||||
result_type out;
|
||||
uint32_t& reg = reinterpret_cast<uint32_t&>(out);
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e4m3x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(reg): "h"(src_packed));
|
||||
|
||||
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<half, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, cutlass::half_t, 2, Round> {
|
||||
using result_element = cutlass::float_e4m3_t;
|
||||
using source_element = cutlass::half_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)
|
||||
uint16_t out;
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f16x2 %0, %1;\n" \
|
||||
"}" \
|
||||
: "=h"(out) : "r"(reinterpret_cast<uint32_t const&>(source)));
|
||||
|
||||
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<half, 2> <= Array<float_e5m2_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<cutlass::half_t, cutlass::float_e5m2_t, 2, Round> {
|
||||
using result_element = cutlass::half_t;
|
||||
using source_element = cutlass::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)
|
||||
result_type out;
|
||||
uint32_t& reg = reinterpret_cast<uint32_t&>(out);
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e5m2x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(reg): "h"(src_packed));
|
||||
|
||||
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_e5m2_t, 2> <= Array<half, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, cutlass::half_t, 2, Round> {
|
||||
using result_element = cutlass::float_e5m2_t;
|
||||
using source_element = cutlass::half_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)
|
||||
uint16_t out;
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.satfinite.e5m2x2.f16x2 %0, %1;\n" \
|
||||
"}" \
|
||||
: "=h"(out) : "r"(reinterpret_cast<uint32_t const&>(source)));
|
||||
|
||||
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 specializations for Array<bfloat16_t, N> <=> Array<float_e4m3_t, N>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Array<bfloat16_t, 2> <= Array<float_e4m3_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<cutlass::bfloat16_t, cutlass::float_e4m3_t, 2, Round> {
|
||||
using result_element = cutlass::bfloat16_t;
|
||||
using source_element = cutlass::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 res_half;
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e4m3x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(res_half): "h"(src_packed));
|
||||
float2 res_float = __half22float2(reinterpret_cast<__half2 &>(res_half));
|
||||
NumericArrayConverter<cutlass::bfloat16_t, float, 2, Round> converter;
|
||||
return converter(reinterpret_cast<Array<float, 2> const&>(res_float));
|
||||
#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<bfloat16_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e4m3_t, cutlass::bfloat16_t, 2, Round> {
|
||||
using result_element = cutlass::float_e4m3_t;
|
||||
using source_element = cutlass::bfloat16_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)
|
||||
NumericArrayConverter<float, cutlass::bfloat16_t, 2, Round> converter;
|
||||
Array<float, 2> res_float = converter(source);
|
||||
uint16_t out;
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 %0, %2, %1;\n" \
|
||||
"}" \
|
||||
: "=h"(out) : "f"(res_float[0]), "f"(res_float[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<bfloat16_t, 2> <= Array<float_e5m2_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<cutlass::bfloat16_t, cutlass::float_e5m2_t, 2, Round> {
|
||||
using result_element = cutlass::bfloat16_t;
|
||||
using source_element = cutlass::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 res_half;
|
||||
uint16_t const& src_packed = reinterpret_cast<uint16_t const&>(source);
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.f16x2.e5m2x2 %0, %1;\n" \
|
||||
"}\n" : "=r"(res_half): "h"(src_packed));
|
||||
float2 res_float = __half22float2(reinterpret_cast<__half2 &>(res_half));
|
||||
NumericArrayConverter<cutlass::bfloat16_t, float, 2, Round> converter;
|
||||
return converter(reinterpret_cast<Array<float, 2> const&>(res_float));
|
||||
#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_e5m2_t, 2> <= Array<bfloat16_t, 2>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<float_e5m2_t, cutlass::bfloat16_t, 2, Round> {
|
||||
using result_element = cutlass::float_e5m2_t;
|
||||
using source_element = cutlass::bfloat16_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)
|
||||
NumericArrayConverter<float, cutlass::bfloat16_t, 2, Round> converter;
|
||||
Array<float, 2> res_float = converter(source);
|
||||
uint16_t out;
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
"cvt.rn.satfinite.e5m2x2.f32 %0, %2, %1;\n" \
|
||||
"}" \
|
||||
: "=h"(out) : "f"(res_float[0]), "f"(res_float[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);
|
||||
}
|
||||
};
|
||||
|
||||
namespace detail {
|
||||
|
||||
/// Special converters that can be used with 4 8-bit elements packed in a register.
|
||||
@@ -2771,86 +3189,6 @@ struct NumericArrayConverter<uint4b_t, int, N, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<int8_t, 8> <= Array<int4b_t, 8>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<int8_t, int4b_t, 8, Round> {
|
||||
|
||||
using result_type = Array<int8_t, 8>;
|
||||
using source_type = Array<int4b_t, 8>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
unsigned const& storage = reinterpret_cast<unsigned const &>(source);
|
||||
unsigned out[2];
|
||||
|
||||
asm volatile(
|
||||
"{ .reg .u32 tmp0, tmp1, tmp2;"
|
||||
"shl.b32 tmp0, %2, 4;"
|
||||
"and.b32 tmp0, tmp0, 0xf0f0f0f0;"
|
||||
"prmt.b32 tmp1, tmp0, tmp0, 0xba98;"
|
||||
"and.b32 tmp1, tmp1, 0xf0f0f0f0;"
|
||||
"shr.u32 tmp0, tmp0, 4;"
|
||||
"or.b32 tmp2, tmp0, tmp1;"
|
||||
"and.b32 tmp0, %2, 0xf0f0f0f0;"
|
||||
"prmt.b32 tmp1, tmp0, tmp0, 0xba98;"
|
||||
"and.b32 tmp1, tmp1, 0xf0f0f0f0;"
|
||||
"shr.u32 tmp0, tmp0, 4;"
|
||||
"or.b32 tmp0, tmp0, tmp1;"
|
||||
"prmt.b32 %0, tmp2, tmp0, 0x5140;"
|
||||
"prmt.b32 %1, tmp2, tmp0, 0x7362;"
|
||||
"}"
|
||||
: "=r"(out[0]), "=r"(out[1])
|
||||
: "r"(storage));
|
||||
|
||||
return reinterpret_cast<result_type const &>(out);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<int8_t> <= Array<int4b_t>
|
||||
template <
|
||||
int N,
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<int8_t, int4b_t, N, Round> {
|
||||
static_assert(!(N % 8), "N must be multiple of 8.");
|
||||
|
||||
using result_type = Array<int8_t, N>;
|
||||
using source_type = Array<int4b_t, N>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
NumericArrayConverter<int8_t, int4b_t, 8, Round> convert_vector_;
|
||||
|
||||
result_type result;
|
||||
|
||||
Array<int8_t, 8> *result_ptr = reinterpret_cast<Array<int8_t, 8> *>(&result);
|
||||
Array<int4b_t, 8> const *source_ptr = reinterpret_cast<Array<int4b_t, 8> const *>(&source);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N / 8; ++i) {
|
||||
result_ptr[i] = convert_vector_(source_ptr[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
#endif // Conditional guards to enable partial specialization for packed integers
|
||||
|
||||
namespace detail {
|
||||
@@ -2942,6 +3280,90 @@ namespace detail {
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
/// Partial specialization for Array<int8_t, 8> <= Array<int4b_t, 8>
|
||||
template <
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<int8_t, int4b_t, 8, Round> {
|
||||
|
||||
using result_type = Array<int8_t, 8>;
|
||||
using source_type = Array<int4b_t, 8>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
unsigned const& storage = reinterpret_cast<unsigned const &>(source);
|
||||
unsigned out[2];
|
||||
|
||||
asm volatile(
|
||||
"{\n"
|
||||
" .reg .u32 tmp0, tmp1, tmp2;\n"
|
||||
" shl.b32 tmp0, %2, 4;\n" // tmp0 = x1x2x3x4x5x6x7__
|
||||
" and.b32 tmp0, tmp0, 0xf0f0f0f0;\n" // tmp0 = x1__x3__x5__x7__
|
||||
" prmt.b32 tmp1, tmp0, tmp0, 0xba98;\n" // tmp1 = s1s3s5s7
|
||||
" and.b32 tmp1, tmp1, 0xf0f0f0f0;\n" // tmp1 = s1__s3__s5__s7__
|
||||
" shr.u32 tmp0, tmp0, 4;\n" // tmp0 = __x1__x3__x5__x7
|
||||
" or.b32 tmp2, tmp0, tmp1;\n" // tmp2 = y1y3y5y7
|
||||
" and.b32 tmp0, %2, 0xf0f0f0f0;\n" // tmp0 = x0__x2__x4__x6__
|
||||
" prmt.b32 tmp1, tmp0, tmp0, 0xba98;\n" // tmp1 = s0s2s4s6
|
||||
" and.b32 tmp1, tmp1, 0xf0f0f0f0;\n" // tmp1 = s0__s2__s4__s6__
|
||||
" shr.u32 tmp0, tmp0, 4;\n" // tmp0 = __x0__x2__x4__x6
|
||||
" or.b32 tmp0, tmp0, tmp1;\n" // tmp0 = y0y2y4y6
|
||||
" prmt.b32 %0, tmp2, tmp0, 0x5140;\n" // %0 = y0y1y2y3
|
||||
" prmt.b32 %1, tmp2, tmp0, 0x7362;\n" // %1 = y4y5y6y7
|
||||
"}\n"
|
||||
: "=r"(out[0]), "=r"(out[1])
|
||||
: "r"(storage));
|
||||
|
||||
return reinterpret_cast<result_type const &>(out);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization for Array<int8_t> <= Array<int4b_t>
|
||||
template <
|
||||
int N,
|
||||
FloatRoundStyle Round
|
||||
>
|
||||
struct NumericArrayConverter<int8_t, int4b_t, N, Round> {
|
||||
static_assert(!(N % 8), "N must be multiple of 8.");
|
||||
|
||||
using result_type = Array<int8_t, N>;
|
||||
using source_type = Array<int4b_t, N>;
|
||||
static FloatRoundStyle const round_style = Round;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static result_type convert(source_type const & source) {
|
||||
|
||||
NumericArrayConverter<int8_t, int4b_t, 8, Round> convert_vector_;
|
||||
|
||||
result_type result;
|
||||
|
||||
Array<int8_t, 8> *result_ptr = reinterpret_cast<Array<int8_t, 8> *>(&result);
|
||||
Array<int4b_t, 8> const *source_ptr = reinterpret_cast<Array<int4b_t, 8> const *>(&source);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N / 8; ++i) {
|
||||
result_ptr[i] = convert_vector_(source_ptr[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
result_type operator()(source_type const &s) const {
|
||||
return convert(s);
|
||||
}
|
||||
};
|
||||
#endif // defined(__CUDA_ARCH__)
|
||||
|
||||
/// Partial specialization for Array<cutlass::float_e4m3_t, N> <= Array<cutlass::int4b_t, N>
|
||||
template <FloatRoundStyle Round, int N>
|
||||
struct NumericArrayConverter<cutlass::float_e4m3_t, cutlass::int4b_t, N, Round> {
|
||||
@@ -3195,6 +3617,16 @@ private:
|
||||
return reinterpret_cast<const uint32_t&>(source);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static int32_t to_int32(source_type_packed_2 const& source) {
|
||||
return static_cast<int32_t>(reinterpret_cast<const int16_t&>(source));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static int32_t to_int32(source_type_packed_4 const& source) {
|
||||
return reinterpret_cast<const int32_t&>(source);
|
||||
}
|
||||
|
||||
template <typename PackedResultType, typename PackedSrcType>
|
||||
CUTLASS_DEVICE
|
||||
static PackedResultType packed_convert(PackedSrcType const &source) {
|
||||
@@ -3206,6 +3638,7 @@ private:
|
||||
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
|
||||
|
||||
PackedResultType r;
|
||||
#if defined __CUDA_ARCH__ && __CUDA_ARCH__ <= 800
|
||||
// View the input as reg
|
||||
uint32_t src_reg = to_reg(source);
|
||||
static constexpr int fp32_base = 0x4B400000;
|
||||
@@ -3223,6 +3656,17 @@ private:
|
||||
result_as_int[ii] += fp32_base;
|
||||
r[ii] -= reinterpret_cast<const float&>(fp32_base);
|
||||
}
|
||||
#else
|
||||
int32_t x = to_int32(source);
|
||||
int32_t t[4];
|
||||
constexpr int32_t mask[4] = {0x00000001, 0x00000100, 0x00010000, 0x01000000};
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ii = 0; ii < PackedResultType::kElements; ++ii) {
|
||||
t[ii] = __dp4a(x, mask[ii], 0);
|
||||
r[ii] = static_cast<float>(t[ii]);
|
||||
}
|
||||
#endif
|
||||
|
||||
return r;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user