CUTLASS 3.0.0 (#786)

* CUTLASS 3.0.0
This commit is contained in:
Vijay Thakkar
2023-01-23 20:55:28 -05:00
committed by GitHub
parent 66d9cddc83
commit 277bd6e537
377 changed files with 76396 additions and 1186 deletions
@@ -62,7 +62,8 @@ template <
typename ElementAccumulator_ = ElementOutput_, ///< Accumulator data type
typename ElementCompute_ = ElementOutput_, ///< Data type used to compute linear combination
ScaleType::Kind Scale = ScaleType::Default, ///< Control Alpha and Beta scaling
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest,
typename ElementSource_ = ElementOutput_
>
class LinearCombination {
public:
@@ -70,6 +71,8 @@ public:
using ElementOutput = ElementOutput_;
using ElementAccumulator = ElementAccumulator_;
using ElementCompute = ElementCompute_;
using ElementC = ElementSource_;
using ElementD = ElementOutput_;
static int const kCount = Count;
static const ScaleType::Kind kScale = Scale;
@@ -78,7 +81,6 @@ public:
using ComputeFragment = Array<ElementCompute, kCount>;
using ParamsBase = LinearCombinationParams;
static FloatRoundStyle const kRound = Round;
/// Host-constructable parameters structure
@@ -89,28 +91,28 @@ public:
ElementCompute const *beta_ptr; ///< pointer to source scalar - if not null, loads it from memory
CUTLASS_HOST_DEVICE
Params():
Params():
ParamsBase(
ElementCompute(1),
ElementCompute(1),
ElementCompute(0)
),
alpha(ElementCompute(1)),
beta(ElementCompute(0)),
alpha_ptr(nullptr),
alpha(ElementCompute(1)),
beta(ElementCompute(0)),
alpha_ptr(nullptr),
beta_ptr(nullptr) { }
CUTLASS_HOST_DEVICE
Params(
ElementCompute alpha,
ElementCompute beta
):
):
ParamsBase(alpha, beta),
alpha(alpha), beta(beta), alpha_ptr(nullptr), beta_ptr(nullptr) { }
alpha(alpha), beta(beta), alpha_ptr(nullptr), beta_ptr(nullptr) { }
CUTLASS_HOST_DEVICE
Params(
ElementCompute alpha
):
):
ParamsBase(alpha, ElementCompute(0)),
alpha(alpha), beta(0), alpha_ptr(nullptr), beta_ptr(nullptr) { }
@@ -118,7 +120,7 @@ public:
Params(
ElementCompute const *alpha_ptr,
ElementCompute const *beta_ptr
):
):
ParamsBase(*alpha_ptr, *beta_ptr),
alpha(0), beta(0), alpha_ptr(alpha_ptr), beta_ptr(beta_ptr) { }
@@ -132,13 +134,13 @@ public:
CUTLASS_HOST_DEVICE
Params(
ParamsBase const& base
): ParamsBase(base), alpha_ptr(nullptr), beta_ptr(nullptr) {
): ParamsBase(base), alpha_ptr(nullptr), beta_ptr(nullptr) {
#if defined(__CUDA_ARCH__)
alpha = reinterpret_cast<ElementCompute const&>(base.alpha_data);
beta = reinterpret_cast<ElementCompute const&>(base.beta_data);
#else
memcpy( alpha, base.alpha_data, sizeof(ElementCompute) );
memcpy( beta, base.alpha_data, sizeof(ElementCompute) );
memcpy( alpha, base.alpha_data, sizeof(ElementCompute) );
memcpy( beta, base.alpha_data, sizeof(ElementCompute) );
#endif
}
};
@@ -184,7 +186,7 @@ public:
/// Computes linear scaling: D = alpha * accumulator + beta * source
CUTLASS_HOST_DEVICE
FragmentOutput operator()(
FragmentAccumulator const &accumulator,
FragmentAccumulator const &accumulator,
FragmentOutput const &source) const {
// Convert source to interal compute numeric type
@@ -236,10 +238,63 @@ public:
ComputeFragment intermediate;
multiplies<ComputeFragment> mul_accumulator;
intermediate = mul_accumulator(alpha_, converted_accumulator); // D = alpha * Accum
intermediate = mul_accumulator(alpha_, converted_accumulator); // D = alpha * Accum
return destination_converter(intermediate);
}
//
// Specializations for scalar (for use with cute::collective::DefaultEpilogue)
//
CUTLASS_HOST_DEVICE
ElementD operator()(ElementAccumulator const accumulator, ElementC const source) const {
// Convert everything to Compute type, do compute, and then store to output type
NumericConverter<ElementCompute, ElementAccumulator, Round> accumulator_converter;
[[maybe_unused]] NumericConverter<ElementCompute, ElementC, Round> source_converter;
NumericConverter<ElementD, ElementCompute, Round> destination_converter;
// Convert to destination numeric type
ElementCompute converted_accumulator = accumulator_converter(accumulator);
if constexpr (Scale == ScaleType::Nothing) {
return destination_converter(converted_accumulator);
}
// Perform binary operations
ElementCompute intermediate;
multiplies<ElementCompute> multiply;
multiply_add<ElementCompute> madd;
if constexpr (Scale == ScaleType::NoBetaScaling) {
intermediate = source_converter(source);
}
else {
intermediate = multiply(beta_, source); // X = beta * C + uniform
}
intermediate = madd(alpha_, converted_accumulator, intermediate); // D = alpha * Accum + X
return destination_converter(intermediate);
}
CUTLASS_HOST_DEVICE
ElementD operator()(ElementAccumulator const accumulator) const {
// Convert everything to Compute type, do compute, and then store to output type
NumericConverter<ElementCompute, ElementAccumulator, Round> accumulator_converter;
NumericConverter<ElementD, ElementCompute, Round> destination_converter;
ElementCompute converted_accumulator = accumulator_converter(accumulator);
// Convert to destination numeric type
if constexpr (Scale == ScaleType::Nothing) {
return destination_converter(converted_accumulator);
}
// Perform binary operations
ElementCompute intermediate;
multiplies<ElementCompute> multiply;
intermediate = multiply(alpha_, accumulator); // D = alpha * Accum
return destination_converter(intermediate);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////