Merge remote-tracking branch 'nccl/master' into develop
Cette révision appartient à :
@@ -15,6 +15,7 @@
|
||||
|
||||
template<typename T>
|
||||
struct FuncNull {
|
||||
__device__ FuncNull(uint64_t opArg=0) {}
|
||||
__device__ T operator()(const T x, const T y) const {
|
||||
return 0;
|
||||
}
|
||||
@@ -22,6 +23,7 @@ struct FuncNull {
|
||||
|
||||
template<typename T>
|
||||
struct FuncSum {
|
||||
__device__ FuncSum(uint64_t opArg=0) {}
|
||||
__device__ T operator()(const T x, const T y) const {
|
||||
return x + y;
|
||||
}
|
||||
@@ -29,6 +31,7 @@ struct FuncSum {
|
||||
|
||||
template<typename T>
|
||||
struct FuncProd {
|
||||
__device__ FuncProd(uint64_t opArg=0) {}
|
||||
__device__ T operator()(const T x, const T y) const {
|
||||
return x * y;
|
||||
}
|
||||
@@ -36,6 +39,7 @@ struct FuncProd {
|
||||
|
||||
template<typename T>
|
||||
struct FuncMax {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
__device__ T operator()(const T x, const T y) const {
|
||||
return (x < y) ? y : x;
|
||||
}
|
||||
@@ -43,6 +47,7 @@ struct FuncMax {
|
||||
|
||||
template<typename T>
|
||||
struct FuncMin {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
__device__ T operator()(const T x, const T y) const {
|
||||
return (x < y) ? x : y;
|
||||
}
|
||||
@@ -53,7 +58,6 @@ struct FuncTraits { // generic implementation for FuncSum,Prod,Min,Max
|
||||
static constexpr bool IsPreOpIdentity = true;
|
||||
static constexpr bool IsPostOpIdentity = true;
|
||||
|
||||
__device__ static Fn make(int rankN) { return Fn(); }
|
||||
template<typename T>
|
||||
__device__ static T preOp(Fn, T x) { return x; }
|
||||
template<typename T>
|
||||
@@ -75,6 +79,7 @@ static __device__ uint32_t addChar4(const uint32_t x, const uint32_t y) {
|
||||
|
||||
template<>
|
||||
struct FuncSum<int8_t> {
|
||||
__device__ FuncSum(uint64_t opArg=0) {}
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
int32_t rv, z=0;
|
||||
@@ -90,6 +95,7 @@ struct FuncSum<int8_t> {
|
||||
};
|
||||
template<>
|
||||
struct FuncSum<uint8_t> {
|
||||
__device__ FuncSum(uint64_t opArg=0) {}
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
int32_t rv, z=0;
|
||||
@@ -119,6 +125,7 @@ static __device__ uint32_t mulChar4(const uint32_t x, const uint32_t y) {
|
||||
|
||||
template<>
|
||||
struct FuncProd<int8_t> {
|
||||
__device__ FuncProd(uint64_t opArg=0) {}
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
return mulChar4(x, y);
|
||||
}
|
||||
@@ -128,6 +135,7 @@ struct FuncProd<int8_t> {
|
||||
};
|
||||
template<>
|
||||
struct FuncProd<uint8_t> {
|
||||
__device__ FuncProd(uint64_t opArg=0) {}
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
return mulChar4(x, y);
|
||||
}
|
||||
@@ -138,6 +146,7 @@ struct FuncProd<uint8_t> {
|
||||
|
||||
template<>
|
||||
struct FuncMax<int8_t> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
union converter { uint32_t storage; char4 a; };
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
@@ -161,6 +170,7 @@ struct FuncMax<int8_t> {
|
||||
};
|
||||
template<>
|
||||
struct FuncMax<uint8_t> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
union converter { uint32_t storage; uchar4 a; };
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
@@ -185,6 +195,7 @@ struct FuncMax<uint8_t> {
|
||||
|
||||
template<>
|
||||
struct FuncMin<int8_t> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
union converter { uint32_t storage; char4 a; };
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
@@ -208,6 +219,7 @@ struct FuncMin<int8_t> {
|
||||
};
|
||||
template<>
|
||||
struct FuncMin<uint8_t> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
union converter { uint32_t storage; uchar4 a; };
|
||||
__device__ uint32_t operator()(const uint32_t x, const uint32_t y) const {
|
||||
#if (__CUDA_ARCH__ >= 300) && (__CUDA_ARCH__ < 500)
|
||||
@@ -232,6 +244,7 @@ struct FuncMin<uint8_t> {
|
||||
|
||||
template<>
|
||||
struct FuncSum<half> {
|
||||
__device__ FuncSum(uint64_t opArg=0) {}
|
||||
__device__ half2 operator()(const half2 x, const half2 y) const {
|
||||
#if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610
|
||||
return __hadd2(x, y);
|
||||
@@ -256,18 +269,16 @@ struct FuncSum<half> {
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct FuncSum<rccl_bfloat16> {
|
||||
__device__ FuncSum(uint64_t opArg=0) {}
|
||||
__device__ rccl_bfloat16 operator()(const rccl_bfloat16 x, const rccl_bfloat16 y) const {
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
return __hadd(x, y);
|
||||
#else
|
||||
return x + y;
|
||||
#endif
|
||||
return (rccl_bfloat16)((float)x + (float)y);
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template<>
|
||||
struct FuncProd<half> {
|
||||
__device__ FuncProd(uint64_t opArg=0) {}
|
||||
__device__ half2 operator()(const half2 x, const half2 y) const {
|
||||
#if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610
|
||||
return __hmul2(x, y);
|
||||
@@ -292,18 +303,16 @@ struct FuncProd<half> {
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct FuncProd<rccl_bfloat16> {
|
||||
__device__ FuncProd(uint64_t opArg=0) {}
|
||||
__device__ rccl_bfloat16 operator()(const rccl_bfloat16 x, const rccl_bfloat16 y) const {
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
return __hmul(x, y);
|
||||
#else
|
||||
return x * y;
|
||||
#endif
|
||||
return (rccl_bfloat16)((float)x * (float)y);
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template<>
|
||||
struct FuncMax<half> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
__device__ half2 operator()(const half2 x, const half2 y) const {
|
||||
float2 fx, fy, fr;
|
||||
fx = __half22float2(x);
|
||||
@@ -324,18 +333,16 @@ struct FuncMax<half> {
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct FuncMax<rccl_bfloat16> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
__device__ rccl_bfloat16 operator()(const rccl_bfloat16 x, const rccl_bfloat16 y) const {
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
return __hmax(x, y);
|
||||
#else
|
||||
return x < y ? y : x;
|
||||
#endif
|
||||
return (float)x < (float)y ? y : x;
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template<>
|
||||
struct FuncMin<half> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
__device__ half2 operator()(const half2 x, const half2 y) const {
|
||||
float2 fx, fy, fr;
|
||||
fx = __half22float2(x);
|
||||
@@ -356,24 +363,23 @@ struct FuncMin<half> {
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct FuncMin<rccl_bfloat16> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
__device__ rccl_bfloat16 operator()(const rccl_bfloat16 x, const rccl_bfloat16 y) const {
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
return __hmin(x, y);
|
||||
#else
|
||||
return x < y ? x : y;
|
||||
#endif
|
||||
return (float)x < (float)y ? x : y;
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
template<>
|
||||
struct FuncMax<float> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
__device__ float operator()(float x, float y) const {
|
||||
return fmaxf(x, y);
|
||||
}
|
||||
};
|
||||
template<>
|
||||
struct FuncMin<float> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
__device__ float operator()(float x, float y) const {
|
||||
return fminf(x, y);
|
||||
}
|
||||
@@ -381,71 +387,98 @@ struct FuncMin<float> {
|
||||
|
||||
template<>
|
||||
struct FuncMax<double> {
|
||||
__device__ FuncMax(uint64_t opArg=0) {}
|
||||
__device__ double operator()(double x, double y) const {
|
||||
return fmax(x, y);
|
||||
}
|
||||
};
|
||||
template<>
|
||||
struct FuncMin<double> {
|
||||
__device__ FuncMin(uint64_t opArg=0) {}
|
||||
__device__ double operator()(double x, double y) const {
|
||||
return fmin(x, y);
|
||||
}
|
||||
};
|
||||
|
||||
template<typename T>
|
||||
struct FuncAvg: FuncSum<T> {
|
||||
static_assert(!std::is_floating_point<T>::value, "Uhoh");
|
||||
struct IsFloatingPoint: std::false_type {};
|
||||
template<>
|
||||
struct IsFloatingPoint<half>: std::true_type {};
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct IsFloatingPoint<rccl_bfloat16>: std::true_type {};
|
||||
#endif
|
||||
template<>
|
||||
struct IsFloatingPoint<float>: std::true_type {};
|
||||
template<>
|
||||
struct IsFloatingPoint<double>: std::true_type {};
|
||||
|
||||
template<typename T, bool IsFloating=IsFloatingPoint<T>::value>
|
||||
struct FuncSumPostDiv;
|
||||
|
||||
template<typename T>
|
||||
struct FuncSumPostDiv<T, /*IsFloating=*/false>: FuncSum<T> {
|
||||
static constexpr bool IsPreOpIdentity = true;
|
||||
static constexpr bool IsPostOpIdentity = false;
|
||||
int n;
|
||||
__device__ FuncSumPostDiv(uint64_t opArg): n(opArg) {}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ T preOp(T x) const { return x; }
|
||||
__device__ T postOp(T x) const { return T(x/n); }
|
||||
};
|
||||
|
||||
template<typename ...Arg>
|
||||
__device__ FuncAvg(int n): n(n) {}
|
||||
template<typename T>
|
||||
struct FuncSumPostDiv<T, /*IsFloating=*/true> {
|
||||
static_assert(sizeof(T)!=sizeof(T), "FuncSumPostDiv is only for implementing ncclAvg on integral types.");
|
||||
};
|
||||
|
||||
__device__ T preOp(T x) const {
|
||||
return x;
|
||||
}
|
||||
__device__ T postOp(T x) const {
|
||||
return T(x/n);
|
||||
}
|
||||
template<typename T>
|
||||
struct FuncPreMulSum: FuncSum<T> { // integral T since all floats are specialized below
|
||||
static constexpr bool IsPreOpIdentity = false;
|
||||
static constexpr bool IsPostOpIdentity = true;
|
||||
T scale;
|
||||
__device__ FuncPreMulSum(uint64_t opArg) { scale = *(T*)&opArg; }
|
||||
// inherits FuncSum::operator()
|
||||
__device__ T preOp(T x) const { return x*scale; }
|
||||
__device__ T postOp(T x) const { return x; }
|
||||
};
|
||||
|
||||
template<>
|
||||
struct FuncAvg<double>: FuncSum<double> {
|
||||
struct FuncPreMulSum<double>: FuncSum<double> {
|
||||
static constexpr bool IsPreOpIdentity = false;
|
||||
static constexpr bool IsPostOpIdentity = true;
|
||||
double rcp;
|
||||
__device__ FuncAvg(int n) {
|
||||
rcp = __drcp_rn(double(n));
|
||||
double scale;
|
||||
__device__ FuncPreMulSum(uint64_t opArg) {
|
||||
scale = *(double*)&opArg;
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ double preOp(double x) const {
|
||||
return IsPreOpIdentity ? x : x*rcp;
|
||||
return IsPreOpIdentity ? x : x*scale;
|
||||
}
|
||||
__device__ double postOp(double x) const {
|
||||
return IsPostOpIdentity ? x : x*rcp;
|
||||
return IsPostOpIdentity ? x : x*scale;
|
||||
}
|
||||
};
|
||||
|
||||
template<>
|
||||
struct FuncAvg<float>: FuncSum<float> {
|
||||
struct FuncPreMulSum<float>: FuncSum<float> {
|
||||
static constexpr bool IsPreOpIdentity = false;
|
||||
static constexpr bool IsPostOpIdentity = true;
|
||||
float rcp;
|
||||
__device__ FuncAvg(int n) {
|
||||
rcp = __frcp_rn(float(n));
|
||||
float scale;
|
||||
__device__ FuncPreMulSum(uint64_t opArg) {
|
||||
scale = *(float*)&opArg;
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ float preOp(float x) const {
|
||||
return IsPreOpIdentity ? x : x*rcp;
|
||||
return IsPreOpIdentity ? x : x*scale;
|
||||
}
|
||||
__device__ float postOp(float x) const {
|
||||
return IsPostOpIdentity ? x : x*rcp;
|
||||
return IsPostOpIdentity ? x : x*scale;
|
||||
}
|
||||
};
|
||||
|
||||
template<>
|
||||
struct FuncAvg<half>: FuncSum<half> {
|
||||
struct FuncPreMulSum<half>: FuncSum<half> {
|
||||
// Change these to switch between all prescale, all postscale, or both by sqrt(N).
|
||||
// Obviously, the only invalid combination is both true. An improvement would be
|
||||
// make this parameterized as a build time setting and passed here through
|
||||
@@ -455,11 +488,8 @@ struct FuncAvg<half>: FuncSum<half> {
|
||||
|
||||
#if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610
|
||||
half2 scale;
|
||||
__device__ FuncAvg(int n) {
|
||||
if (!IsPreOpIdentity && !IsPostOpIdentity)
|
||||
scale.x = __float2half(__frsqrt_rn(float(n)));
|
||||
else
|
||||
scale.x = __float2half(__frcp_rn(float(n)));
|
||||
__device__ FuncPreMulSum(uint64_t opArg) {
|
||||
scale.x = *(half*)&opArg;
|
||||
scale.y = scale.x;
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
@@ -477,11 +507,8 @@ struct FuncAvg<half>: FuncSum<half> {
|
||||
}
|
||||
#else
|
||||
float scale;
|
||||
__device__ FuncAvg(int n) {
|
||||
if (!IsPreOpIdentity && !IsPostOpIdentity)
|
||||
scale = __frsqrt_rn(float(n));
|
||||
else
|
||||
scale = __frcp_rn(float(n));
|
||||
__device__ FuncPreMulSum(uint64_t opArg) {
|
||||
scale = __half2float(*(half*)&opArg);
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ half preOp(half x) const {
|
||||
@@ -515,64 +542,54 @@ struct FuncAvg<half>: FuncSum<half> {
|
||||
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
template<>
|
||||
struct FuncAvg<rccl_bfloat16>: FuncSum<rccl_bfloat16> {
|
||||
struct FuncPreMulSum<rccl_bfloat16>: FuncSum<rccl_bfloat16> {
|
||||
// Change these to switch between all prescale, all postscale, or both by sqrt(N).
|
||||
// Obviously, the only invalid combination is both true. An improvement would be
|
||||
// make this parameterized as a build time setting and passed here through
|
||||
// preprocessor definitions.
|
||||
static constexpr bool IsPreOpIdentity = true;
|
||||
static constexpr bool IsPostOpIdentity = false;
|
||||
static constexpr bool IsPreOpIdentity = false;
|
||||
static constexpr bool IsPostOpIdentity = true;
|
||||
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
__device__ FuncAvg(int n) {
|
||||
if (!IsPreOpIdentity && !IsPostOpIdentity)
|
||||
scale.x = __float2bfloat16(__frsqrt_rn(float(n)));
|
||||
else
|
||||
scale.x = __float2bfloat16(__frcp_rn(float(n)));
|
||||
scale.y = scale.x;
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ rccl_bfloat16 preOp(rccl_bfloat16 x) const {
|
||||
return IsPreOpIdentity ? x : __hmul(x, scale.x);
|
||||
}
|
||||
__device__ rccl_bfloat16 postOp(rccl_bfloat16 x) const {
|
||||
return IsPostOpIdentity ? x : __hmul(x, scale.x);
|
||||
}
|
||||
#else
|
||||
float scale;
|
||||
__device__ FuncAvg(int n) {
|
||||
if (!IsPreOpIdentity && !IsPostOpIdentity)
|
||||
scale = __frsqrt_rn(float(n));
|
||||
else
|
||||
scale = __frcp_rn(float(n));
|
||||
__device__ FuncPreMulSum(uint64_t opArg) {
|
||||
scale = *(rccl_bfloat16*)&opArg;
|
||||
}
|
||||
// inherits FuncSum::operator()
|
||||
__device__ rccl_bfloat16 preOp(rccl_bfloat16 x) const {
|
||||
return IsPreOpIdentity ? x : (rccl_bfloat16)(x*scale);
|
||||
return IsPreOpIdentity ? x : (rccl_bfloat16)((float)x*scale);
|
||||
}
|
||||
__device__ rccl_bfloat16 postOp(rccl_bfloat16 x) const {
|
||||
return IsPostOpIdentity ? x : (rccl_bfloat16)(x*scale);
|
||||
return IsPostOpIdentity ? x : (rccl_bfloat16)((float)x*scale);
|
||||
}
|
||||
#endif
|
||||
};
|
||||
#endif
|
||||
|
||||
template<typename T>
|
||||
struct FuncTraits<FuncAvg<T>> {
|
||||
static constexpr bool IsPreOpIdentity = FuncAvg<T>::IsPreOpIdentity;
|
||||
static constexpr bool IsPostOpIdentity = FuncAvg<T>::IsPostOpIdentity;
|
||||
struct FuncTraits<FuncPreMulSum<T>> {
|
||||
static constexpr bool IsPreOpIdentity = FuncPreMulSum<T>::IsPreOpIdentity;
|
||||
static constexpr bool IsPostOpIdentity = FuncPreMulSum<T>::IsPostOpIdentity;
|
||||
|
||||
__device__ static FuncAvg<T> make(int rankN) {
|
||||
return FuncAvg<T>(rankN);
|
||||
}
|
||||
template<typename U>
|
||||
__device__ static U preOp(FuncAvg<T> fn, U x) {
|
||||
__device__ static U preOp(FuncPreMulSum<T> fn, U x) {
|
||||
return fn.preOp(x);
|
||||
}
|
||||
template<typename U>
|
||||
__device__ static U postOp(FuncAvg<T> fn, U x) {
|
||||
__device__ static U postOp(FuncPreMulSum<T> fn, U x) {
|
||||
return fn.postOp(x);
|
||||
}
|
||||
};
|
||||
template<typename T>
|
||||
struct FuncTraits<FuncSumPostDiv<T>> {
|
||||
static constexpr bool IsPreOpIdentity = FuncSumPostDiv<T>::IsPreOpIdentity;
|
||||
static constexpr bool IsPostOpIdentity = FuncSumPostDiv<T>::IsPostOpIdentity;
|
||||
|
||||
template<typename U>
|
||||
__device__ static U preOp(FuncSumPostDiv<T> fn, U x) {
|
||||
return fn.preOp(x);
|
||||
}
|
||||
template<typename U>
|
||||
__device__ static U postOp(FuncSumPostDiv<T> fn, U x) {
|
||||
return fn.postOp(x);
|
||||
}
|
||||
};
|
||||
#endif // REDUCE_KERNEL_H_
|
||||
|
||||
Référencer dans un nouveau ticket
Bloquer un utilisateur