From 1cda2f52b6a3b96a5035c049519a988f8e3bccfa Mon Sep 17 00:00:00 2001 From: Wenkai Du Date: Fri, 15 Nov 2019 13:46:03 -0800 Subject: [PATCH] Add bf16 support in rccl-tests --- src/Makefile | 3 +- src/common.cu | 26 ++++- src/common.h | 4 + src/rccl_bfloat16.h | 253 ++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 283 insertions(+), 3 deletions(-) create mode 100644 src/rccl_bfloat16.h diff --git a/src/Makefile b/src/Makefile index 78470b8f48..157a351e5c 100644 --- a/src/Makefile +++ b/src/Makefile @@ -15,11 +15,10 @@ HIPCC = $(ROCM_HOME)/hip/bin/hipcc CXX = $(HIPCC) -HIPCUFLAGS := +HIPCUFLAGS := -std=c++14 HIPCUFLAGS += -I$(ROCM_HOME)/include HIPCUFLAGS += -I$(ROCM_HOME)/include/rccl HIPCUFLAGS += -I$(ROCM_HOME)/hip/include/hip -HIPCUFLAGS += -I$(ROCM_HOME)/hiprand/include LDFLAGS := -L$(ROCM_HOME)/lib -lhsa-runtime64 -lrt HIPLDFLAGS := $(CUSTOM_RCCL_LIB) -L$(ROCM_HOME)/lib -lhsa-runtime64 -lrt diff --git a/src/common.cu b/src/common.cu index 5bf78eeebe..07ebcd90a3 100644 --- a/src/common.cu +++ b/src/common.cu @@ -7,6 +7,7 @@ ************************************************************************/ #include "hip/hip_runtime.h" +#include "rccl_bfloat16.h" #include "common.h" #include #include @@ -16,8 +17,13 @@ #include #if NCCL_MAJOR >= 2 +#if RCCL_BFLOAT16 == 1 +ncclDataType_t test_types[ncclNumTypes] = {ncclInt8, ncclUint8, ncclInt32, ncclUint32, ncclInt64, ncclUint64, ncclHalf, ncclFloat, ncclDouble, ncclBfloat16}; +const char *test_typenames[ncclNumTypes] = {"int8", "uint8", "int32", "uint32", "int64", "uint64", "half", "float", "double", "bf16"}; +#else ncclDataType_t test_types[ncclNumTypes] = {ncclInt8, ncclUint8, ncclInt32, ncclUint32, ncclInt64, ncclUint64, ncclHalf, ncclFloat, ncclDouble}; const char *test_typenames[ncclNumTypes] = {"int8", "uint8", "int32", "uint32", "int64", "uint64", "half", "float", "double"}; +#endif #else ncclDataType_t test_types[ncclNumTypes] = {ncclChar, ncclInt, ncclHalf, ncclFloat, ncclDouble, ncclInt64, ncclUint64}; const char *test_typenames[ncclNumTypes] = {"char", "int", "half", "float", "double", "int64", "uint64"}; @@ -78,6 +84,9 @@ double DeltaMaxValue(ncclDataType_t type) { #endif case ncclInt64: case ncclUint64: return 1e-200; +#if NCCL_MAJOR >= 2 && RCCL_BFLOAT16 == 1 + case ncclBfloat16: return 1e-2; +#endif } return 1e-200; } @@ -155,6 +164,10 @@ testResult_t CheckDelta(void* expected, void* results, size_t count, ncclDataTyp case ncclInt64: case ncclUint64: hipLaunchKernelGGL((deltaKern), dim3(1), dim3(512), 0, 0, results, expected, count, devmax); break; +#if NCCL_MAJOR >= 2 && RCCL_BFLOAT16 == 1 + case ncclBfloat16: + hipLaunchKernelGGL((deltaKern), dim3(1), dim3(512), 0, 0, results, expected, count, devmax); break; +#endif } HIPCHECK(hipDeviceSynchronize()); return testSuccess; @@ -181,6 +194,10 @@ template<> __device__ half testValue(const size_t offset, const int rep, const int rank) { return __float2half(testValue(offset, rep, rank)); } +template<> +__device__ rccl_bfloat16 testValue(const size_t offset, const int rep, const int rank) { + return (float)testValue(offset, rep, rank); +} // Operations template @@ -220,7 +237,11 @@ typedef void(*redInitKern_t)(void* data, const size_t N, const size_t offset, co static redInitKern_t const redInitDataKerns[ncclNumOps*ncclNumTypes] = { #if NCCL_MAJOR >= 2 +#if RCCL_BFLOAT16 == 1 + OPS(int8_t), OPS(uint8_t), OPS(int32_t), OPS(uint32_t), OPS(int64_t), OPS(uint64_t), OPS(half), OPS(float), OPS(double), OPS(rccl_bfloat16) +#else OPS(int8_t), OPS(uint8_t), OPS(int32_t), OPS(uint32_t), OPS(int64_t), OPS(uint64_t), OPS(half), OPS(float), OPS(double) +#endif #else OPS(char), OPS(int32_t), OPS(half), OPS(float), OPS(double), OPS(int64_t), OPS(uint64_t) #endif @@ -251,7 +272,10 @@ static initDataKern_t const initDataKerns[ncclNumTypes] = { InitDataKernel, InitDataKernel< half>, InitDataKernel< float>, - InitDataKernel< double> + InitDataKernel< double>, +#if RCCL_BFLOAT16 == 1 + InitDataKernel +#endif #else InitDataKernel< char>, InitDataKernel< int32_t>, diff --git a/src/common.h b/src/common.h index dd98d547df..54f216c9c5 100644 --- a/src/common.h +++ b/src/common.h @@ -195,6 +195,10 @@ static size_t wordSize(ncclDataType_t type) { case ncclDouble: //case ncclFloat64: return 8; +#if NCCL_MAJOR >= 2 && RCCL_BFLOAT16 == 1 + case ncclBfloat16: + return 2; +#endif default: return 0; } } diff --git a/src/rccl_bfloat16.h b/src/rccl_bfloat16.h new file mode 100644 index 0000000000..06b053a626 --- /dev/null +++ b/src/rccl_bfloat16.h @@ -0,0 +1,253 @@ +/** + * MIT License + * + * Copyright 2019 Advanced Micro Devices, Inc. All rights reserved. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE + * SOFTWARE. + */ + +/*!\file + * \brief rccl_bfloat16.h provides struct for rccl_bfloat16 typedef + */ + +#ifndef _RCCL_BFLOAT16_H_ +#define _RCCL_BFLOAT16_H_ + +#if __cplusplus < 201402L || (!defined(__HCC__) && !defined(__HIPCC__)) + +// If this is a C compiler, C++ compiler below C++14, or a host-only compiler, we only +// include a minimal definition of rccl_bfloat16 + +#include +/*! \brief Struct to represent a 16 bit brain floating point number. */ +typedef struct +{ + uint16_t data; +} rccl_bfloat16; + +#else // __cplusplus < 201402L || (!defined(__HCC__) && !defined(__HIPCC__)) + +#include +#include +#include +#include +#include +#include + +struct rccl_bfloat16 +{ + uint16_t data; + + __host__ __device__ rccl_bfloat16() = default; + + // round upper 16 bits of IEEE float to convert to bfloat16 + explicit constexpr __host__ __device__ rccl_bfloat16(float f) + : data(float_to_bfloat16(f)) + { + } + + // zero extend lower 16 bits of bfloat16 to convert to IEEE float + constexpr __host__ __device__ operator float() const + { + union + { + uint32_t int32; + float fp32; + } u = {uint32_t(data) << 16}; + return u.fp32; + } + +private: + static constexpr __host__ __device__ uint16_t float_to_bfloat16(float f) + { + union + { + float fp32; + uint32_t int32; + } u = {f}; + if(~u.int32 & 0x7f800000) + { + // When the exponent bits are not all 1s, then the value is zero, normal, + // or subnormal. We round the bfloat16 mantissa up by adding 0x7FFF, plus + // 1 if the least significant bit of the bfloat16 mantissa is 1 (odd). + // This causes the bfloat16's mantissa to be incremented by 1 if the 16 + // least significant bits of the float mantissa are greater than 0x8000, + // or if they are equal to 0x8000 and the least significant bit of the + // bfloat16 mantissa is 1 (odd). This causes it to be rounded to even when + // the lower 16 bits are exactly 0x8000. If the bfloat16 mantissa already + // has the value 0x7f, then incrementing it causes it to become 0x00 and + // the exponent is incremented by one, which is the next higher FP value + // to the unrounded bfloat16 value. When the bfloat16 value is subnormal + // with an exponent of 0x00 and a mantissa of 0x7F, it may be rounded up + // to a normal value with an exponent of 0x01 and a mantissa of 0x00. + // When the bfloat16 value has an exponent of 0xFE and a mantissa of 0x7F, + // incrementing it causes it to become an exponent of 0xFF and a mantissa + // of 0x00, which is Inf, the next higher value to the unrounded value. + u.int32 += 0x7fff + ((u.int32 >> 16) & 1); // Round to nearest, round to even + } + else if(u.int32 & 0xffff) + { + // When all of the exponent bits are 1, the value is Inf or NaN. + // Inf is indicated by a zero mantissa. NaN is indicated by any nonzero + // mantissa bit. Quiet NaN is indicated by the most significant mantissa + // bit being 1. Signaling NaN is indicated by the most significant + // mantissa bit being 0 but some other bit(s) being 1. If any of the + // lower 16 bits of the mantissa are 1, we set the least significant bit + // of the bfloat16 mantissa, in order to preserve signaling NaN in case + // the bloat16's mantissa bits are all 0. + u.int32 |= 0x10000; // Preserve signaling NaN + } + return uint16_t(u.int32 >> 16); + } +}; + +typedef struct +{ + uint16_t data; +} rccl_bfloat16_public; + +static_assert(std::is_standard_layout{}, + "rccl_bfloat16 is not a standard layout type, and thus is " + "incompatible with C."); + +static_assert(std::is_trivial{}, + "rccl_bfloat16 is not a trivial type, and thus is " + "incompatible with C."); + +static_assert(sizeof(rccl_bfloat16) == sizeof(rccl_bfloat16_public) + && offsetof(rccl_bfloat16, data) == offsetof(rccl_bfloat16_public, data), + "internal rccl_bfloat16 does not match public rccl_bfloat16"); + +inline std::ostream& operator<<(std::ostream& os, const rccl_bfloat16& bf16) +{ + return os << float(bf16); +} +constexpr __host__ __device__ rccl_bfloat16 operator+(rccl_bfloat16 a) +{ + return a; +} +constexpr __host__ __device__ rccl_bfloat16 operator-(rccl_bfloat16 a) +{ + a.data ^= 0x8000; + return a; +} +constexpr __host__ __device__ rccl_bfloat16 operator+(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return rccl_bfloat16(float(a) + float(b)); +} +constexpr __host__ __device__ rccl_bfloat16 operator-(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return rccl_bfloat16(float(a) - float(b)); +} +constexpr __host__ __device__ rccl_bfloat16 operator*(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return rccl_bfloat16(float(a) * float(b)); +} +constexpr __host__ __device__ rccl_bfloat16 operator/(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return rccl_bfloat16(float(a) / float(b)); +} +constexpr __host__ __device__ bool operator<(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return float(a) < float(b); +} +constexpr __host__ __device__ bool operator==(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return float(a) == float(b); +} +constexpr __host__ __device__ bool operator>(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return b < a; +} +constexpr __host__ __device__ bool operator<=(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return !(a > b); +} +constexpr __host__ __device__ bool operator!=(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return !(a == b); +} +constexpr __host__ __device__ bool operator>=(rccl_bfloat16 a, rccl_bfloat16 b) +{ + return !(a < b); +} +constexpr __host__ __device__ rccl_bfloat16& operator+=(rccl_bfloat16& a, rccl_bfloat16 b) +{ + return a = a + b; +} +constexpr __host__ __device__ rccl_bfloat16& operator-=(rccl_bfloat16& a, rccl_bfloat16 b) +{ + return a = a - b; +} +constexpr __host__ __device__ rccl_bfloat16& operator*=(rccl_bfloat16& a, rccl_bfloat16 b) +{ + return a = a * b; +} +constexpr __host__ __device__ rccl_bfloat16& operator/=(rccl_bfloat16& a, rccl_bfloat16 b) +{ + return a = a / b; +} +constexpr __host__ __device__ rccl_bfloat16& operator++(rccl_bfloat16& a) +{ + return a += rccl_bfloat16(1.0f); +} +constexpr __host__ __device__ rccl_bfloat16& operator--(rccl_bfloat16& a) +{ + return a -= rccl_bfloat16(1.0f); +} +constexpr __host__ __device__ rccl_bfloat16 operator++(rccl_bfloat16& a, int) +{ + rccl_bfloat16 orig = a; + ++a; + return orig; +} +constexpr __host__ __device__ rccl_bfloat16 operator--(rccl_bfloat16& a, int) +{ + rccl_bfloat16 orig = a; + --a; + return orig; +} + +namespace std +{ + constexpr __host__ __device__ bool isinf(rccl_bfloat16 a) + { + return !(~a.data & 0x7f80) && !(a.data & 0x7f); + } + constexpr __host__ __device__ bool isnan(rccl_bfloat16 a) + { + return !(~a.data & 0x7f80) && +(a.data & 0x7f); + } + constexpr __host__ __device__ bool iszero(rccl_bfloat16 a) + { + return !(a.data & 0x7fff); + } + inline rccl_bfloat16 sin(rccl_bfloat16 a) + { + return rccl_bfloat16(sinf(float(a))); + } + inline rccl_bfloat16 cos(rccl_bfloat16 a) + { + return rccl_bfloat16(cosf(float(a))); + } +} + +#endif // __cplusplus < 201402L || (!defined(__HCC__) && !defined(__HIPCC__)) + +#endif // _RCCL_BFLOAT16_H_