Enable fp8 support (#1101)

* initial checkin

* resolve cr comments

* resolve the build issue

* fix the data correctless issue

* update fp8 header file and update the unit test for fp8 support

* remove fp16 from fp8 headers

* fix ut issue and catch up the latest code from develop

* udate according to cr comments

* update ut according to cr comments

* update num floats for each SumPostDiv from 4 to 6

* update fp8 header file name

* fix the typo
This commit is contained in:
Andy li
2024-03-09 07:17:53 +08:00
committed by GitHub
parent ff951e607d
commit 6777e65c1d
29 changed files with 1243 additions and 48 deletions
+3 -1
View File
@@ -414,7 +414,9 @@ __global__ void MSCCL_KERNEL_ENTRY_NAME(devredop, type, Simple, fullOps)(struct
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, half, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, float, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, double, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps)
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_float8, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat8, fullOps)
#define MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_NOFLOAT(devredop, fullOps) \
MSCCL_IMPL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, int8_t, fullOps) \
+7 -2
View File
@@ -1,5 +1,6 @@
/*************************************************************************
* Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved.
* Modifications Copyright (c) Microsoft Corporation. Licensed under the MIT License.
*
* See LICENSE.txt for license information
************************************************************************/
@@ -63,9 +64,13 @@ ncclResult_t ncclLaunchOneRank(void* dst, void const* src, size_t nElts, struct
case ncclInt64: kernel = (void const*)&oneRankReduce<FuncPreMulSum<int64_t>>; break;
case ncclUint64: kernel = (void const*)&oneRankReduce<FuncPreMulSum<uint64_t>>; break;
case ncclFloat16: kernel = (void const*)&oneRankReduce<FuncPreMulSum<half>>; break;
#if defined(RCCL_BFLOAT16)
#if defined(RCCL_BFLOAT16)
case ncclBfloat16: kernel = (void const*)&oneRankReduce<FuncPreMulSum<rccl_bfloat16>>; break;
#endif
#endif
#if defined(RCCL_FLOAT8)
case ncclFp8E4M3: kernel = (void const*)&oneRankReduce<FuncPreMulSum<rccl_float8>>; break;
case ncclFp8E5M2: kernel = (void const*)&oneRankReduce<FuncPreMulSum<rccl_bfloat8>>; break;
#endif
case ncclFloat32: kernel = (void const*)&oneRankReduce<FuncPreMulSum<float>>; break;
case ncclFloat64: kernel = (void const*)&oneRankReduce<FuncPreMulSum<double>>; break;
default: return ncclInvalidArgument;
+74
View File
@@ -1,6 +1,7 @@
/*************************************************************************
* Copyright (c) 2015-2021, NVIDIA CORPORATION. All rights reserved.
* Modifications Copyright (c) 2019-2021 Advanced Micro Devices, Inc. All rights reserved.
* Modifications Copyright (c) Microsoft Corporation. Licensed under the MIT License.
*
* See LICENSE.txt for license information
************************************************************************/
@@ -13,6 +14,8 @@
#include <limits>
#include <type_traits>
#include "rccl_float8.h"
template<typename T>
struct IsFloatingPoint: std::false_type {};
template<>
@@ -21,6 +24,12 @@ struct IsFloatingPoint<half>: std::true_type {};
template<>
struct IsFloatingPoint<rccl_bfloat16>: std::true_type {};
#endif
#if defined(RCCL_FLOAT8)
template<>
struct IsFloatingPoint<rccl_float8>: std::true_type {};
template<>
struct IsFloatingPoint<rccl_bfloat8>: std::true_type {};
#endif
template<>
struct IsFloatingPoint<float>: std::true_type {};
template<>
@@ -254,6 +263,15 @@ SPECIALIZE_REDUCE(FuncMinMax, half, 1, half, fn.isMinNotMax ? __hmin(x, y) : __h
#endif
#endif
#if defined(RCCL_FLOAT8)
SPECIALIZE_REDUCE(FuncSum, rccl_float8, 1, rccl_float8, rccl_float8(float(x) + float(y)))
SPECIALIZE_REDUCE(FuncProd, rccl_float8, 1, rccl_float8, rccl_float8(float(x) * float(y)))
SPECIALIZE_REDUCE(FuncMinMax, rccl_float8, 1, rccl_float8, rccl_float8(fn.isMinNotMax ? fminf(float(x), float(y)) : fmaxf(float(x), float(y))))
SPECIALIZE_REDUCE(FuncSum, rccl_bfloat8, 1, rccl_bfloat8, rccl_bfloat8(float(x) + float(y)))
SPECIALIZE_REDUCE(FuncProd, rccl_bfloat8, 1, rccl_bfloat8, rccl_bfloat8(float(x) * float(y)))
SPECIALIZE_REDUCE(FuncMinMax, rccl_bfloat8, 1, rccl_bfloat8, rccl_bfloat8(fn.isMinNotMax ? fminf(float(x), float(y)) : fmaxf(float(x), float(y))))
#endif
#undef SPECIALIZE_REDUCE
////////////////////////////////////////////////////////////////////////////////
@@ -389,6 +407,38 @@ struct FuncPreMulSum<half> {
};
#endif
#if defined(RCCL_FLOAT8)
template<>
struct FuncPreMulSum<rccl_float8> {
// 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.
using EltType = rccl_float8;
float scalar;
__device__ FuncPreMulSum(uint64_t opArg=0) {
union { uint64_t u64; rccl_float8 val; };
u64 = opArg;
scalar = (float)(val);
}
};
template<>
struct FuncPreMulSum<rccl_bfloat8> {
// 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.
using EltType = rccl_bfloat8;
float scalar;
__device__ FuncPreMulSum(uint64_t opArg=0) {
union { uint64_t u64; rccl_bfloat8 val; };
u64 = opArg;
scalar = (float)(val);
}
};
#endif
template<typename T>
struct Apply_Reduce<FuncPreMulSum<T>, /*EltPerPack=*/1> {
__device__ static BytePack<sizeof(T)> reduce(FuncPreMulSum<T> fn, BytePack<sizeof(T)> a, BytePack<sizeof(T)> b) {
@@ -456,6 +506,30 @@ struct Apply_PreOp<FuncPreMulSum<half>, /*EltPerPack=*/1> {
#endif
#endif
#if defined(RCCL_FLOAT8)
template<>
struct Apply_PreOp<FuncPreMulSum<rccl_float8>, /*EltPerPack=*/1> {
static constexpr bool IsIdentity = false;
__device__ static BytePack<sizeof(rccl_float8)> preOp(
FuncPreMulSum<rccl_float8> fn, BytePack<sizeof(rccl_float8)> a
) {
return toPack<rccl_float8>(rccl_float8(float(fromPack<rccl_float8>(a)) * float(fn.scalar)));
}
};
template<>
struct Apply_PreOp<FuncPreMulSum<rccl_bfloat8>, /*EltPerPack=*/1> {
static constexpr bool IsIdentity = false;
__device__ static BytePack<sizeof(rccl_bfloat8)> preOp(
FuncPreMulSum<rccl_bfloat8> fn, BytePack<sizeof(rccl_bfloat8)> a
) {
return toPack<rccl_bfloat8>(rccl_bfloat8(float(fromPack<rccl_bfloat8>(a)) * float(fn.scalar)));
}
};
#endif
////////////////////////////////////////////////////////////////////////////////
// FuncSumPostDiv
+20 -5
View File
@@ -1,6 +1,7 @@
/*************************************************************************
* Copyright (c) 2017-2022, NVIDIA CORPORATION. All rights reserved.
* Modifications Copyright (c) 2019-2023 Advanced Micro Devices, Inc. All rights reserved.
* Modifications Copyright (c) Microsoft Corporation. Licensed under the MIT License.
*
* See LICENSE.txt for license information
************************************************************************/
@@ -1556,9 +1557,13 @@ static ncclResult_t hostToDevRedOp(
half f16;
float f32;
double f64;
#if defined(RCCL_BFLOAT16)
rccl_bfloat16 bf16;
#endif
#if defined(RCCL_BFLOAT16)
rccl_bfloat16 bf16;
#endif
#if defined(RCCL_FLOAT8)
rccl_float8 fp8_e4m3;
rccl_bfloat8 fp8_e5m2;
#endif
void *ptr;
};
u64 = 0;
@@ -1594,12 +1599,22 @@ static ncclResult_t hostToDevRedOp(
opFull->op = ncclDevPreMulSum;
f16 = __float2half(float(1.0/comm->nRanks)); // __double2half not supported pre CUDA 11.x
break;
#if defined(RCCL_BFLOAT16)
#if defined(RCCL_BFLOAT16)
case ncclBfloat16:
opFull->op = ncclDevPreMulSum;
bf16 = (rccl_bfloat16)(float(1.0/comm->nRanks));
break;
#endif
#endif
#if defined(RCCL_FLOAT8)
case ncclFp8E4M3:
opFull->op = ncclDevPreMulSum;
fp8_e4m3 = static_cast<rccl_float8>(float(1.0/comm->nRanks));
break;
case ncclFp8E5M2:
opFull->op = ncclDevPreMulSum;
fp8_e5m2 = static_cast<rccl_bfloat8>(float(1.0/comm->nRanks));
break;
#endif
case ncclFloat32:
opFull->op = ncclDevPreMulSum;
f32 = float(1.0/comm->nRanks);
+7 -2
View File
@@ -1,6 +1,7 @@
/*************************************************************************
* Copyright (c) 2017-2022, NVIDIA CORPORATION. All rights reserved.
* Modifications Copyright (c) 2019-2022 Advanced Micro Devices, Inc. All rights reserved.
* Modifications Copyright (c) Microsoft Corporation. Licensed under the MIT License.
*
* See LICENSE.txt for license information
************************************************************************/
@@ -29,11 +30,15 @@ inline int ncclTypeSize(ncclDataType_t type) {
switch (type) {
case ncclInt8:
case ncclUint8:
#if defined(RCCL_FLOAT8)
case ncclFp8E4M3:
case ncclFp8E5M2:
#endif
return 1;
case ncclFloat16:
#if defined(RCCL_BFLOAT16)
#if defined(RCCL_BFLOAT16)
case ncclBfloat16:
#endif
#endif
return 2;
case ncclInt32:
case ncclUint32:
+14 -8
View File
@@ -1,6 +1,7 @@
/*************************************************************************
* Copyright (c) 2015-2022, NVIDIA CORPORATION. All rights reserved.
* Modifications Copyright (c) 2019-2022 Advanced Micro Devices, Inc. All rights reserved.
* Modifications Copyright (c) Microsoft Corporation. Licensed under the MIT License.
*
* See LICENSE.txt for license information
************************************************************************/
@@ -9,6 +10,7 @@
#define NCCL_DEVICE_H_
#include "nccl.h"
#include "rccl_float8.h"
#include "rccl_bfloat16.h"
#include "nccl_common.h"
#include "align.h"
@@ -502,9 +504,13 @@ inline bool ncclNvlsSupported(int devRedOp, int type) {
case ncclInt64:
case ncclUint64:
case ncclFloat16:
#if defined(RCCL_BFLOAT16)
#if defined(RCCL_BFLOAT16)
case ncclBfloat16:
#endif
#endif
#if defined(RCCL_FLOAT8)
case ncclFp8E4M3:
case ncclFp8E5M2:
#endif
return devRedOp == ncclDevSum || devRedOp == ncclDevMinMax;
case ncclFloat:
case ncclDouble:
@@ -530,10 +536,10 @@ inline int ncclDevFuncId(int coll, int devRedOp, int type, int algo, int proto)
// <all_algos> / <all_protos> / <all_redops> / <all_types>
if (coll == ncclFuncAllReduce) {
row += (((algo * NCCL_NUM_PROTOCOLS + proto) * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * (algo * NCCL_NUM_PROTOCOLS + proto);
row += (((algo * NCCL_NUM_PROTOCOLS + proto) * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 6 * (algo * NCCL_NUM_PROTOCOLS + proto);
goto have_row;
}
row += (NCCL_NUM_ALGORITHMS - 2) * NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
row += (NCCL_NUM_ALGORITHMS - 2) * NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 6);
// RING / SIMPLE / Sum / int8_t
if (coll == ncclFuncAllToAllPivot) goto have_row;
@@ -548,17 +554,17 @@ inline int ncclDevFuncId(int coll, int devRedOp, int type, int algo, int proto)
// RING / <all_protos> / <all_redops> / <all_types>
if (coll == ncclFuncReduce) {
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * proto;
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 6 * proto;
goto have_row;
}
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 6);
// RING / <all_protos> / <all_redops> / <all_types>
if (coll == ncclFuncReduceScatter) {
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * proto;
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 6 * proto;
goto have_row;
}
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 6);
// RING / SIMPLE / Sum / int8_t
if (coll == ncclFuncSendRecv) goto have_row;
+3 -1
View File
@@ -26,7 +26,9 @@ __global__ void MSCCL_KERNEL_ENTRY_NAME(devredop, type, proto, fullOps)(struct n
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, half, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, float, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, double, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps)
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_float8, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, rccl_bfloat8, fullOps)
#define MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_NOFLOAT(devredop, fullOps) \
MSCCL_DECL_KERNEL_ENTRY_FUNC_DEVREDOP_TYPE(devredop, int8_t, fullOps) \
+3 -3
View File
@@ -15,9 +15,9 @@ typedef void (*ncclDebugLogger_t)(ncclDebugLogLevel level, unsigned long flags,
#define NCCL_NUM_FUNCTIONS 5 // Send/Recv and AllToAllPivot not included for now
typedef enum { ncclFuncBroadcast, ncclFuncReduce, ncclFuncAllGather, ncclFuncReduceScatter, ncclFuncAllReduce, ncclFuncSendRecv, ncclFuncSend, ncclFuncRecv, ncclFuncAllToAllPivot, ncclNumFuncs} ncclFunc_t;
#define FUNC_INDEX_P2P 835
#define FUNC_INDEX_ALLTOALL_PIVOT 555
#define FUNC_INDEX_TOTAL 846
#define FUNC_INDEX_P2P 979
#define FUNC_INDEX_ALLTOALL_PIVOT 651
#define FUNC_INDEX_TOTAL 992
#define NCCL_NUM_ALGORITHMS 6 // Tree/Ring/CollNet*
#define NCCL_ALGO_UNDEF -1
File diff suppressed because it is too large Load Diff
+19 -1
View File
@@ -226,6 +226,10 @@ static ncclResult_t hostToDevRedOp(
#if defined(RCCL_BFLOAT16)
rccl_bfloat16 bf16;
#endif
#if defined(RCCL_FLOAT8)
rccl_float8 fp8_e4m3;
rccl_bfloat8 fp8_e5m2;
#endif
float f32;
double f64;
void *ptr;
@@ -269,6 +273,16 @@ static ncclResult_t hostToDevRedOp(
bf16 = (rccl_bfloat16)(float(1.0/comm->nRanks));
break;
#endif
#if defined(RCCL_FLOAT8)
case ncclFp8E4M3:
opFull->op = ncclDevPreMulSum;
fp8_e4m3 = (rccl_float8)(float(1.0/comm->nRanks));
break;
case ncclFp8E5M2:
opFull->op = ncclDevPreMulSum;
fp8_e5m2 = (rccl_bfloat8)(float(1.0/comm->nRanks));
break;
#endif
case ncclFloat32:
opFull->op = ncclDevPreMulSum;
f32 = float(1.0/comm->nRanks);
@@ -315,7 +329,9 @@ static ncclResult_t hostToDevRedOp(
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, half, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, float, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, double, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps)
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, rccl_float8, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, rccl_bfloat8, fullOps)
#define MSCCL_KERNEL_ENTRY_DEVREDOP_NOFLOAT(devredop, fullOps) \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, int8_t, fullOps), \
@@ -327,6 +343,8 @@ static ncclResult_t hostToDevRedOp(
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL()
#define MSCCL_KERNEL_ENTRY() \
+7
View File
@@ -21,6 +21,7 @@
#define NCCL_VERSION(X,Y,Z) (((X) <= 2 && (Y) <= 8) ? (X) * 1000 + (Y) * 100 + (Z) : (X) * 10000 + (Y) * 100 + (Z))
#define RCCL_BFLOAT16 1
#define RCCL_FLOAT8 1
#define RCCL_GATHER_SCATTER 1
#define RCCL_ALLTOALLV 1
@@ -362,7 +363,13 @@ typedef enum { ncclInt8 = 0, ncclChar = 0,
ncclFloat32 = 7, ncclFloat = 7,
ncclFloat64 = 8, ncclDouble = 8,
ncclBfloat16 = 9,
#if defined(RCCL_FLOAT8)
ncclFp8E4M3 = 10,
ncclFp8E5M2 = 11,
ncclNumTypes = 12 } ncclDataType_t;
#else
ncclNumTypes = 10 } ncclDataType_t;
#endif
/*! @} */
/*! @defgroup rccl_api_custom_redop Custom Reduction Operator