/************************************************************************* * Copyright (c) 2015-2019, NVIDIA CORPORATION. All rights reserved. * * See LICENSE.txt for license information ************************************************************************/ #ifndef NCCL_REDUCE_KERNEL_H_ #define NCCL_REDUCE_KERNEL_H_ #include "common_kernel.h" #include template struct FuncNull { __device__ T operator()(const T x, const T y) const { return 0; } }; template struct FuncSum { __device__ T operator()(const T x, const T y) const { return x + y; } }; template struct FuncProd { __device__ T operator()(const T x, const T y) const { return x * y; } }; template struct FuncMax { __device__ T operator()(const T x, const T y) const { return (x < y) ? y : x; } }; template struct FuncMin { __device__ T operator()(const T x, const T y) const { return (x < y) ? x : y; } }; #define MASK0 0x00ff00ff #define MASK1 0xff00ff00 static __device__ uint32_t addChar4(const uint32_t x, const uint32_t y) { /* This can be used both for signed and unsigned 8-bit addition */ const uint32_t x0 = x & MASK0; const uint32_t x1 = x & MASK1; const uint32_t y0 = y & MASK0; const uint32_t y1 = y & MASK1; const uint32_t r0 = (x0+y0); const uint32_t r1 = (x1+y1); return (r0 & MASK0) | (r1 & MASK1); } template<> struct FuncSum { __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; asm("vadd4.s32.s32.s32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else return addChar4(x, y); #endif } __device__ int8_t operator()(const int8_t x, const int8_t y) const { return x+y; } }; template<> struct FuncSum { __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; asm("vadd4.u32.u32.u32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else return addChar4(x, y); #endif } __device__ uint8_t operator()(const uint8_t x, const uint8_t y) const { return x+y; } }; static __device__ uint32_t mulChar4(const uint32_t x, const uint32_t y) { /* This can be used both for signed and unsigned 8-bit multiplication */ union converter { uint32_t storage; char4 a; }; converter cx, cy, cr; cx.storage = x; cy.storage = y; cr.a.x = cx.a.x * cy.a.x; cr.a.y = cx.a.y * cy.a.y; cr.a.z = cx.a.z * cy.a.z; cr.a.w = cx.a.w * cy.a.w; return cr.storage; } template<> struct FuncProd { __device__ uint32_t operator()(const uint32_t x, const uint32_t y) const { return mulChar4(x, y); } __device__ int8_t operator()(const int8_t x, const int8_t y) const { return x*y; } }; template<> struct FuncProd { __device__ uint32_t operator()(const uint32_t x, const uint32_t y) const { return mulChar4(x, y); } __device__ uint8_t operator()(const uint8_t x, const uint8_t y) const { return x*y; } }; template<> struct FuncMax { 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) int32_t rv, z=0; asm("vmax4.s32.s32.s32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else converter cx, cy, cr; cx.storage = x; cy.storage = y; cr.a.x = max(cx.a.x, cy.a.x); cr.a.y = max(cx.a.y, cy.a.y); cr.a.z = max(cx.a.z, cy.a.z); cr.a.w = max(cx.a.w, cy.a.w); return cr.storage; #endif } __device__ int8_t operator()(const int8_t x, const int8_t y) const { return (x>y) ? x : y; } }; template<> struct FuncMax { 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) int32_t rv, z=0; asm("vmax4.u32.u32.u32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else converter cx, cy, cr; cx.storage = x; cy.storage = y; cr.a.x = max(cx.a.x, cy.a.x); cr.a.y = max(cx.a.y, cy.a.y); cr.a.z = max(cx.a.z, cy.a.z); cr.a.w = max(cx.a.w, cy.a.w); return cr.storage; #endif } __device__ uint8_t operator()(const uint8_t x, const uint8_t y) const { return (x>y) ? x : y; } }; template<> struct FuncMin { 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) int32_t rv, z=0; asm("vmin4.s32.s32.s32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else converter cx, cy, cr; cx.storage = x; cy.storage = y; cr.a.x = min(cx.a.x, cy.a.x); cr.a.y = min(cx.a.y, cy.a.y); cr.a.z = min(cx.a.z, cy.a.z); cr.a.w = min(cx.a.w, cy.a.w); return cr.storage; #endif } __device__ int8_t operator()(const int8_t x, const int8_t y) const { return (x struct FuncMin { 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) int32_t rv, z=0; asm("vmin4.u32.u32.u32 %0, %1, %2, %3;" : "=r"(rv) : "r"(x), "r"(y), "r"(z)); return rv; #else converter cx, cy, cr; cx.storage = x; cy.storage = y; cr.a.x = min(cx.a.x, cy.a.x); cr.a.y = min(cx.a.y, cy.a.y); cr.a.z = min(cx.a.z, cy.a.z); cr.a.w = min(cx.a.w, cy.a.w); return cr.storage; #endif } __device__ uint8_t operator()(const uint8_t x, const uint8_t y) const { return (x struct FuncSum { __device__ half2 operator()(const half2 x, const half2 y) const { #if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610 return __hadd2(x, y); #else float2 fx, fy, fr; fx = __half22float2(x); fy = __half22float2(y); fr.x = fx.x + fy.x; fr.y = fx.y + fy.y; return __float22half2_rn(fr); #endif } __device__ half operator()(const half x, const half y) const { #if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610 return __hadd(x, y); #else return __float2half( __half2float(x) + __half2float(y) ); #endif } }; template<> struct FuncProd { __device__ half2 operator()(const half2 x, const half2 y) const { #if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610 return __hmul2(x, y); #else float2 fx, fy, fr; fx = __half22float2(x); fy = __half22float2(y); fr.x = fx.x * fy.x; fr.y = fx.y * fy.y; return __float22half2_rn(fr); #endif } __device__ half operator()(const half x, const half y) const { #if __CUDA_ARCH__ >= 530 && __CUDA_ARCH__ != 610 return __hmul(x, y); #else return __float2half( __half2float(x) * __half2float(y) ); #endif } }; template<> struct FuncMax { __device__ half2 operator()(const half2 x, const half2 y) const { float2 fx, fy, fr; fx = __half22float2(x); fy = __half22float2(y); fr.x = fmaxf(fx.x, fy.x); fr.y = fmaxf(fx.y, fy.y); return __float22half2_rn(fr); } __device__ half operator()(const half x, const half y) const { float fx, fy, fm; fx = __half2float(x); fy = __half2float(y); fm = fmaxf(fx, fy); return __float2half(fm); } }; template<> struct FuncMin { __device__ half2 operator()(const half2 x, const half2 y) const { float2 fx, fy, fr; fx = __half22float2(x); fy = __half22float2(y); fr.x = fminf(fx.x, fy.x); fr.y = fminf(fx.y, fy.y); return __float22half2_rn(fr); } __device__ half operator()(const half x, const half y) const { float fx, fy, fm; fx = __half2float(x); fy = __half2float(y); fm = fminf(fx, fy); return __float2half(fm); } }; #endif // REDUCE_KERNEL_H_