Updates for CUTLASS 3.5.0 (#1468)

This commit is contained in:
Vijay Thakkar
2024-04-11 21:33:40 -04:00
committed by GitHub
parent a40e08e9d5
commit 7d49e6c7e2
171 changed files with 7526 additions and 1888 deletions

View File

@@ -33,16 +33,6 @@
\brief Boost-like numeric conversion operator for CUTLASS numeric types
*/
/*
Note: CUTLASS 3x increases the host compiler requirements to C++17. However, certain
existing integrations of CUTLASS require C++11 host compilers.
Until this requirement can be lifted, certain headers with this annotation are required
to be remain consistent with C++11 syntax.
C++11 compatibility is enforced by this unit test: `cutlass_test_unit_core_cpp11`.
*/
#pragma once
#if !defined(__CUDACC_RTC__)
@@ -848,18 +838,21 @@ struct NumericArrayConverter<cutlass::half_t, float, 2, FloatRoundStyle::round_t
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
Array<cutlass::half_t, 2> result;
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 530)
Array<cutlass::half_t, 2> result;
reinterpret_cast<__half2 &>(result) = __float22half2_rn(reinterpret_cast<float2 const &>(source));
return result;
#else
NumericConverter<cutlass::half_t, float, round_style> convert_;
// NOTE: cutlass::Array<half, N> is NOT an aggregate type and
// below `{}` does NOT conduct zero initialization. Below `{}` will
// conduct default initialization (calling default ctr). We use this syntax
// to resolve compiler warning on uninitialized member variable.
Array<cutlass::half_t, 2> result{};
result[0] = convert_(source[0]);
result[1] = convert_(source[1]);
return result;
#endif
return result;
}
CUTLASS_HOST_DEVICE
@@ -879,17 +872,19 @@ struct NumericArrayConverter<float, cutlass::half_t, 2, Round> {
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
Array<float, 2> result;
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 530)
reinterpret_cast<float2 &>(result) = __half22float2(reinterpret_cast<__half2 const &>(source));
float2 result2 = __half22float2(reinterpret_cast<__half2 const &>(source));
return {
float{result2.x},
float{result2.y}
};
#else
NumericConverter<float, cutlass::half_t, round_style> convert_;
result[0] = convert_(source[0]);
result[1] = convert_(source[1]);
return {
convert_(source[0]),
convert_(source[1])
};
#endif
return result;
}
CUTLASS_HOST_DEVICE
@@ -1482,7 +1477,7 @@ struct NumericArrayConverterPacked4Element {
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]));
}
@@ -2306,14 +2301,69 @@ template <
struct NumericArrayConverter<float_e5m2_t, cutlass::float_e5m2_t, N, Round> :
public PackedNumericArrayConverter<float_e5m2_t, cutlass::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.
template <
FloatRoundStyle Round
>
struct NumericArrayConverter<int8_t, float, 1, Round> {
using result_type = Array<int8_t, 1>;
using source_type = Array<float, 1>;
static FloatRoundStyle const round_style = Round;
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
// Convert to int to int8_t
NumericConverter<int8_t, float, Round> destination_converter;
result_type result;
result[0] = destination_converter(source[0]);
return result;
}
CUTLASS_HOST_DEVICE
result_type operator()(source_type const &s) const {
return convert(s);
}
};
// To convert a FP32 to Int that has less than 32 bits, we need to convert it to int32 first.
template <
typename T,
int N,
FloatRoundStyle Round
>
struct NumericArrayFP32ToIntConverter {
using result_type = Array<T, N>;
using source_type = Array<float, N>;
static FloatRoundStyle const round_style = Round;
static_assert(platform::numeric_limits<T>::is_integer, "the dest type has to be int.");
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
// Convert float to int
Array<int32_t, N> temporary;
NumericArrayConverter<int32_t, float, N, Round> compute_converter;
temporary = compute_converter(source);
// Convert to int to int8_t
NumericArrayConverter<T, int32_t, N, Round> destination_converter;
return destination_converter(temporary);
}
CUTLASS_HOST_DEVICE
result_type operator()(source_type const &s) const {
return convert(s);
}
};
template <
int N,
FloatRoundStyle Round
@@ -2322,19 +2372,74 @@ struct NumericArrayConverter<int8_t, float, N, Round> {
using result_type = Array<int8_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) {
// Convert float to int
Array<int32_t, N> temporary;
NumericArrayFP32ToIntConverter<int8_t, N, Round> converter;
return converter(source);
}
NumericArrayConverter<int, float, N, Round> compute_converter;
temporary = compute_converter(source);
CUTLASS_HOST_DEVICE
result_type operator()(source_type const &s) const {
return convert(s);
}
};
// Convert to int to int8_t
NumericArrayConverter<int8_t, int32_t, N, Round> destination_converter;
return destination_converter(temporary);
template <
int N,
FloatRoundStyle Round
>
struct NumericArrayConverter<uint8_t, float, N, Round> {
using result_type = Array<uint8_t, N>;
using source_type = Array<float, N>;
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
NumericArrayFP32ToIntConverter<uint8_t, N, Round> converter;
return converter(source);
}
CUTLASS_HOST_DEVICE
result_type operator()(source_type const &s) const {
return convert(s);
}
};
template <
int N,
FloatRoundStyle Round
>
struct NumericArrayConverter<int4b_t, float, N, Round> {
using result_type = Array<int4b_t, N>;
using source_type = Array<float, N>;
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
NumericArrayFP32ToIntConverter<int4b_t, N, Round> converter;
return converter(source);
}
CUTLASS_HOST_DEVICE
result_type operator()(source_type const &s) const {
return convert(s);
}
};
template <
int N,
FloatRoundStyle Round
>
struct NumericArrayConverter<uint4b_t, float, N, Round> {
using result_type = Array<uint4b_t, N>;
using source_type = Array<float, N>;
CUTLASS_HOST_DEVICE
static result_type convert(source_type const & source) {
NumericArrayFP32ToIntConverter<uint4b_t, N, Round> converter;
return converter(source);
}
CUTLASS_HOST_DEVICE
@@ -2508,7 +2613,7 @@ namespace detail {
template <int Offset, size_t ParentWidth, typename ArrayConverter>
CUTLASS_DEVICE
static void convert_helper(
typename ArrayConverter::result_type& result,
typename ArrayConverter::result_type& result,
typename ArrayConverter::source_type const& source) {
using ElementRes = typename ArrayConverter::result_type::Element;
@@ -2530,14 +2635,14 @@ namespace detail {
static void convert_helper(typename ArrayConverter::result_type& result, typename ArrayConverter::source_type const& source) {
static_assert(sizeof...(OtherVectorArrays) % 2 == 0, "Vector converters must come in {dst, src} pairs");
static_assert(ResultVectorArray::kElements == SourceVectorArray::kElements, "Vector converters must have the same vector width");
static_assert(cutlass::platform::is_same<typename ArrayConverter::result_type::Element, typename ResultVectorArray::Element>::value,
static_assert(cutlass::platform::is_same<typename ArrayConverter::result_type::Element, typename ResultVectorArray::Element>::value,
"ResultVectorArray must have the same type ArrayConverter::result_type");
static_assert(cutlass::platform::is_same<typename ArrayConverter::source_type::Element, typename SourceVectorArray::Element>::value,
static_assert(cutlass::platform::is_same<typename ArrayConverter::source_type::Element, typename SourceVectorArray::Element>::value,
"SourceVectorArray must have the same type ArrayConverter::result_type");
static_assert(Offset >= 0 && Offset <= ArrayConverter::result_type::kElements, "Offset must be between 0 and N");
static_assert(ParentWidth == 0 || ParentWidth > ResultVectorArray::kElements, "Vector arrays must be given in decreasing order of width");
constexpr int vector_width = ResultVectorArray::kElements;
static_assert(ispow2(vector_width), "Vector width must be a power of 2");
@@ -2569,8 +2674,8 @@ namespace detail {
public:
/*
A method to convert vectors of elements using the packed_convert method of the converter.
A method to convert vectors of elements using the packed_convert method of the converter.
Converters using this class must implement packed convert and support 1 or more vector conversions.
*/
template <typename ArrayConverter, typename ResultVectorArray, typename SourceVectorArray, typename... OtherVectorArrays>
@@ -2651,7 +2756,7 @@ private:
uint32_t final_prmt_idx = final_prmt_base | sign;
// This uses a look up table to convert packed int4s to packed fp8s, using the int4 value
// as the index to prmt.
// as the index to prmt.
// It first select both the positive and negative candidates, then uses the sign bit to
// select the correct candidate.
asm volatile(
@@ -2675,8 +2780,8 @@ public:
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4>(result, source);
return result;
@@ -2684,7 +2789,7 @@ public:
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -2771,7 +2876,7 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_8>::value &&
platform::is_same<PackedResultType, result_type_packed_8>::value),
"Invalid PackedSrcType/PackedResultType must be 1, 2, 4 or 8 to use private convert dispatch.");
// Hold output FP16s in reg. We need 1 reg for every 2 elements
PackedResultType r;
@@ -2788,23 +2893,23 @@ private:
return r;
}
friend class detail::VectorizedConverter;
friend class detail::VectorizedConverter;
public:
CUTLASS_DEVICE
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -2844,7 +2949,7 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
PackedResultType r;
// View the input as reg
uint32_t src_reg = to_reg(source);
@@ -2875,15 +2980,15 @@ public:
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -2923,12 +3028,12 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
PackedResultType r;
// View the input as reg
uint32_t src_reg = to_reg(source);
// __byte_perm simulates the add.u32 0x4B000000 to every u8 element of u8x4 source and stores
// __byte_perm simulates the add.u32 0x4B000000 to every u8 element of u8x4 source and stores
// the result in r (without introducing extra cvt.u32.u8 instruction)
uint32_t const prmt_indices[4] = {0x7650, 0x7651, 0x7652, 0x7653};
uint32_t* result_as_int = reinterpret_cast<uint32_t*>(&r);
@@ -2948,15 +3053,15 @@ public:
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3010,10 +3115,10 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_8>::value &&
platform::is_same<PackedResultType, result_type_packed_8>::value),
"Invalid PackedSrcType/PackedResultType must be 2, 4 or 8 to use private convert dispatch.");
// Hold output FP16s in reg. We need 1 reg for every 2 elements
using RegArray = cutlass::AlignedArray<uint32_t, PackedResultType::kElements / 2, sizeof(PackedResultType)>;
RegArray r;
RegArray r;
// View the input as reg
uint32_t src_reg = to_reg(source);
@@ -3034,7 +3139,7 @@ private:
" prmt.b32 %0, %1, %2, %3;\n"
"}\n"
: "=r"(r[ii])
: "r"(src_reg), "n"(0), "r"(prmt_indices[ii]));
: "r"(src_reg), "n"(0), "r"(prmt_indices[ii]));
}
// The below XOR does the following:
@@ -3057,7 +3162,7 @@ private:
" lop3.b32 %0, %0, %1, %2, %3;\n"
"}\n"
: "+r"(r[ii])
: "n"(and_mask), "n"(xor_mask), "n"(immLut));
: "n"(and_mask), "n"(xor_mask), "n"(immLut));
}
// We will issue 2 hfmas that do the following:
@@ -3087,23 +3192,23 @@ private:
return reinterpret_cast<PackedResultType&>(r);
}
friend class detail::VectorizedConverter;
friend class detail::VectorizedConverter;
public:
CUTLASS_DEVICE
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3145,10 +3250,10 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
// Hold output FP16s in reg. We need 1 reg for every 2 elements
using RegArray = cutlass::AlignedArray<uint32_t, PackedResultType::kElements / 2, sizeof(PackedResultType)>;
RegArray r;
RegArray r;
#if 0 // Scalar conversion (Please keep this code for reference for vectorized version below)
auto result = reinterpret_cast<PackedResultType&>(r);
@@ -3176,18 +3281,18 @@ private:
// 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
// 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;
// If we want F = ((a & b) ^ c) then set kImmLut = (0xF0 & 0xCC) ^ 0xAA
static constexpr uint32_t kImmLut = (0xF0 & 0xCC) ^ 0xAA;
for (int ii = 0; ii < RegArray::kElements; ++ii) {
// The bit-wise operation executed below is `r[ii] = (r[ii] & 0x03FF03FF) ^ 0x66006600;`
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n" :
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n" :
"=r"(r[ii]) : "r"(r[ii]), "n"(0x03FF03FF), "n"(0x66006600), "n"(kImmLut));
}
@@ -3209,14 +3314,14 @@ public:
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3256,11 +3361,11 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
// Hold output FP16s in reg. We need 1 reg for every 2 elements
using RegArray = cutlass::AlignedArray<uint32_t, PackedResultType::kElements / 2, sizeof(PackedResultType)>;
RegArray r;
RegArray r;
// View the input as reg
uint32_t src_reg = to_reg(source);
uint32_t const prmt_indices[2] = {0x5150, 0x5352};
@@ -3289,15 +3394,15 @@ public:
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3352,10 +3457,10 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_8>::value &&
platform::is_same<PackedResultType, result_type_packed_8>::value),
"Invalid PackedSrcType/PackedResultType must be 2, 4 or 8 to use private convert dispatch.");
// Hold output FP16s in reg. We need 1 reg for every 2 elements
using RegArray = cutlass::AlignedArray<uint32_t, PackedResultType::kElements / 2, sizeof(PackedResultType)>;
RegArray r;
RegArray r;
// View the input as reg
uint32_t src_reg = to_reg(source);
@@ -3371,7 +3476,7 @@ private:
" prmt.b32 %0, %1, %2, %3;\n"
"}\n"
: "=r"(r[ii])
: "r"(src_reg), "r"(src_reg_shifted), "r"(prmt_indices[ii]));
: "r"(src_reg), "r"(src_reg_shifted), "r"(prmt_indices[ii]));
}
// The below XOR does the following:
@@ -3390,7 +3495,7 @@ private:
" lop3.b32 %0, %0, %1, %2, %3;\n"
"}\n"
: "+r"(r[ii])
: "n"(and_mask), "n"(xor_mask), "n"(immLut));
: "n"(and_mask), "n"(xor_mask), "n"(immLut));
}
// We will issue 2 bfmas that do the following:
@@ -3400,7 +3505,7 @@ private:
// This is the BF16 {136, 136} represented as an integer.
static constexpr uint32_t bias_rep = 0x43084308;
const __nv_bfloat162& bias = reinterpret_cast<const __nv_bfloat162&>(bias_rep);
CUTLASS_PRAGMA_UNROLL
for (int ii = 0; ii < RegArray::kElements; ++ii) {
__nv_bfloat162& bf16x2_val = reinterpret_cast<__nv_bfloat162&>(r[ii]);
@@ -3410,23 +3515,23 @@ private:
return reinterpret_cast<PackedResultType&>(r);
}
friend class detail::VectorizedConverter;
friend class detail::VectorizedConverter;
public:
CUTLASS_DEVICE
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_8, source_type_packed_8,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3466,7 +3571,7 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
NumericArrayConverter<float, int8_t, PackedResultType::kElements, Round> convert_int8_to_f32;
Array<float, PackedResultType::kElements> tmp = convert_int8_to_f32(source);
NumericArrayConverter<cutlass::bfloat16_t, float, PackedResultType::kElements, Round> convert_f32_to_bf16;
@@ -3481,15 +3586,15 @@ public:
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3529,7 +3634,7 @@ private:
(platform::is_same<PackedSrcType, source_type_packed_4>::value &&
platform::is_same<PackedResultType, result_type_packed_4>::value),
"Invalid PackedSrcType/PackedResultType must be 2 or 4 to use private convert dispatch.");
NumericArrayConverter<float, uint8_t, PackedResultType::kElements, Round> convert_uint8_to_f32;
Array<float, PackedResultType::kElements> tmp = convert_uint8_to_f32(source);
NumericArrayConverter<cutlass::bfloat16_t, float, PackedResultType::kElements, Round> convert_f32_to_bf16_;
@@ -3543,15 +3648,15 @@ public:
static result_type convert(source_type const &source) {
result_type result;
using ConverterType = NumericArrayConverter<typename result_type::Element, typename source_type::Element, N, Round>;
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
detail::VectorizedConverter::convert<ConverterType,
result_type_packed_4, source_type_packed_4,
result_type_packed_2, source_type_packed_2>(result, source);
return result;
}
CUTLASS_DEVICE
result_type operator()(source_type const &s) const {
result_type operator()(source_type const &s) const {
return convert(s);
}
};
@@ -3705,7 +3810,7 @@ struct PackPredicates {
int word_idx = (i / kWordSize);
int bit_idx = (i % kWordSize);
uint8_t mask = ((predicates[i] ? 1u : 0u) << bit_idx);
uint8_t mask = static_cast<uint8_t>((predicates[i] ? 1u : 0u) << bit_idx);
bytes[word_idx] = (bytes[word_idx] | mask);
}
return packed;