@@ -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);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user