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
+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);