Merge remote-tracking branch 'nccl/master' into develop
This commit is contained in:
@@ -112,10 +112,10 @@ union PtrUnion
|
||||
int64_t* I8; // ncclInt64
|
||||
uint64_t* U8; // ncclUint64
|
||||
__half* F2; // ncclFloat16
|
||||
rccl_float8* F1; // ncclFp8E4M3
|
||||
rccl_float8* F1; // ncclFloat8e4m3
|
||||
float* F4; // ncclFloat32
|
||||
double* F8; // ncclFloat64
|
||||
rccl_bfloat8* B1; // ncclFp8E5M2
|
||||
rccl_bfloat8* B1; // ncclFloat8e5m2
|
||||
hip_bfloat16* B2; // ncclBfloat16
|
||||
|
||||
constexpr PtrUnion() : ptr(nullptr) {}
|
||||
@@ -176,8 +176,8 @@ std::string DataTypeToName(ncclDataType_t const dataType)
|
||||
case ncclFloat32: return "Float32";
|
||||
case ncclFloat64: return "Float64";
|
||||
case ncclBfloat16: return "Bfloat16";
|
||||
case ncclFp8E4M3: return "Fp8E4M3";
|
||||
case ncclFp8E5M2: return "Fp8E5M2";
|
||||
case ncclFloat8e4m3: return "Fp8E4M3";
|
||||
case ncclFloat8e5m2: return "Fp8E5M2";
|
||||
default:
|
||||
printf("Unsupported datatype (%d)\n", dataType);
|
||||
exit(0);
|
||||
@@ -197,8 +197,8 @@ size_t DataTypeToBytes(ncclDataType_t const dataType)
|
||||
case ncclFloat32: return 4;
|
||||
case ncclFloat64: return 8;
|
||||
case ncclBfloat16: return 2;
|
||||
case ncclFp8E4M3: return 1;
|
||||
case ncclFp8E5M2: return 1;
|
||||
case ncclFloat8e4m3: return 1;
|
||||
case ncclFloat8e5m2: return 1;
|
||||
default:
|
||||
printf("Unsupported datatype (%s)\n", DataTypeToName(dataType).c_str());
|
||||
exit(0);
|
||||
@@ -239,11 +239,11 @@ void SetPtr(PtrUnion& ptrUnion, ncclDataType_t const dataType, int const idx, in
|
||||
case ncclUint32: ptrUnion.U4[idx] = valueI; break;
|
||||
case ncclInt64: ptrUnion.I8[idx] = valueI; break;
|
||||
case ncclUint64: ptrUnion.U8[idx] = valueI; break;
|
||||
case ncclFp8E4M3: ptrUnion.F1[idx] = rccl_float8(valueF); break;
|
||||
case ncclFloat8e4m3: ptrUnion.F1[idx] = rccl_float8(valueF); break;
|
||||
case ncclFloat16: ptrUnion.F2[idx] = __float2half(static_cast<float>(valueF)); break;
|
||||
case ncclFloat32: ptrUnion.F4[idx] = valueF; break;
|
||||
case ncclFloat64: ptrUnion.F8[idx] = valueF; break;
|
||||
case ncclFp8E5M2: ptrUnion.B1[idx] = rccl_bfloat8(valueF); break;
|
||||
case ncclFloat8e5m2: ptrUnion.B1[idx] = rccl_bfloat8(valueF); break;
|
||||
case ncclBfloat16: ptrUnion.B2[idx] = hip_bfloat16(static_cast<float>(valueF)); break;
|
||||
default:
|
||||
printf("Unsupported datatype (%s)\n", DataTypeToName(dataType).c_str());
|
||||
@@ -265,11 +265,11 @@ bool IsEqual(PtrUnion const& actual, PtrUnion const& expected, ncclDataType_t co
|
||||
case ncclUint32: isMatch = (actual.U4[idx] == expected.U4[idx]); break;
|
||||
case ncclInt64: isMatch = (actual.I8[idx] == expected.I8[idx]); break;
|
||||
case ncclUint64: isMatch = (actual.U8[idx] == expected.U8[idx]); break;
|
||||
case ncclFp8E4M3: isMatch = (fabs(float(actual.F1[idx]) - float(expected.F1[idx])) < 9e-2); break;
|
||||
case ncclFloat8e4m3: isMatch = (fabs(float(actual.F1[idx]) - float(expected.F1[idx])) < 9e-2); break;
|
||||
case ncclFloat16: isMatch = (fabs(__half2float(actual.F2[idx]) - __half2float(expected.F2[idx])) < 9e-2); break;
|
||||
case ncclFloat32: isMatch = (fabs(actual.F4[idx] - expected.F4[idx]) < 1e-5); break;
|
||||
case ncclFloat64: isMatch = (fabs(actual.F8[idx] - expected.F8[idx]) < 1e-12); break;
|
||||
case ncclFp8E5M2: isMatch = (fabs(float(actual.B1[idx]) - float(expected.B1[idx])) < 9e-2); break;
|
||||
case ncclFloat8e5m2: isMatch = (fabs(float(actual.B1[idx]) - float(expected.B1[idx])) < 9e-2); break;
|
||||
case ncclBfloat16: isMatch = (fabs((float)actual.B2[idx] - (float)expected.B2[idx]) < 9e-2); break;
|
||||
default:
|
||||
printf("Unsupported datatype (%s)\n", DataTypeToName(dataType).c_str());
|
||||
@@ -290,7 +290,7 @@ bool IsEqual(PtrUnion const& actual, PtrUnion const& expected, ncclDataType_t co
|
||||
printf("[Error Rank = %d] Expected output: %ld. Actual output: %ld at index %lu\n", globalRank, expected.I8[idx], actual.I8[idx], idx); break;
|
||||
case ncclUint64:
|
||||
printf("[Error Rank = %d] Expected output: %lu. Actual output: %lu at index %lu\n", globalRank, expected.U8[idx], actual.U8[idx], idx); break;
|
||||
case ncclFp8E4M3:
|
||||
case ncclFloat8e4m3:
|
||||
printf("[Error Rank = %d] Expected output: %f. Actual output: %f at index %lu\n", globalRank, (float)expected.F1[idx], (float)actual.F1[idx], idx); break;
|
||||
case ncclFloat16:
|
||||
printf("[Error Rank = %d] Expected output: %f. Actual output: %f at index %lu\n", globalRank, __half2float(expected.F2[idx]), __half2float(actual.F2[idx]), idx); break;
|
||||
@@ -298,7 +298,7 @@ bool IsEqual(PtrUnion const& actual, PtrUnion const& expected, ncclDataType_t co
|
||||
printf("[Error Rank = %d] Expected output: %f. Actual output: %f at index %lu\n", globalRank, expected.F4[idx], actual.F4[idx], idx); break;
|
||||
case ncclFloat64:
|
||||
printf("[Error Rank = %d] Expected output: %lf. Actual output: %lf at index %lu\n", globalRank, expected.F8[idx], actual.F8[idx], idx); break;
|
||||
case ncclFp8E5M2:
|
||||
case ncclFloat8e5m2:
|
||||
printf("[Error Rank = %d] Expected output: %f. Actual output: %f at index %lu\n", globalRank, (float)expected.B1[idx], (float)actual.B1[idx], idx); break;
|
||||
case ncclBfloat16:
|
||||
printf("[Error Rank = %d] Expected output: %f. Actual output: %f at index %lu\n", globalRank, (float)expected.B2[idx], (float)actual.B2[idx], idx); break;
|
||||
@@ -340,11 +340,11 @@ void Reduce(PtrUnion& ptrUnion, PtrUnion const& otherPtrUnion, size_t const numE
|
||||
case ncclUint32: ptrUnion.U4[idx] = ReduceOp(op, ptrUnion.U4[idx], otherPtrUnion.U4[idx]); break;
|
||||
case ncclInt64: ptrUnion.I8[idx] = ReduceOp(op, ptrUnion.I8[idx], otherPtrUnion.I8[idx]); break;
|
||||
case ncclUint64: ptrUnion.U8[idx] = ReduceOp(op, ptrUnion.U8[idx], otherPtrUnion.U8[idx]); break;
|
||||
case ncclFp8E4M3: ptrUnion.F1[idx] = rccl_float8(ReduceOp(op, float(ptrUnion.F1[idx]), float(otherPtrUnion.F1[idx]))); break;
|
||||
case ncclFloat8e4m3: ptrUnion.F1[idx] = rccl_float8(ReduceOp(op, float(ptrUnion.F1[idx]), float(otherPtrUnion.F1[idx]))); break;
|
||||
case ncclFloat16: ptrUnion.F2[idx] = __float2half(ReduceOp(op, __half2float(ptrUnion.F2[idx]), __half2float(otherPtrUnion.F2[idx]))); break;
|
||||
case ncclFloat32: ptrUnion.F4[idx] = ReduceOp(op, ptrUnion.F4[idx], otherPtrUnion.F4[idx]); break;
|
||||
case ncclFloat64: ptrUnion.F8[idx] = ReduceOp(op, ptrUnion.F8[idx], otherPtrUnion.F8[idx]); break;
|
||||
case ncclFp8E5M2: ptrUnion.B1[idx] = rccl_bfloat8(ReduceOp(op, float(ptrUnion.B1[idx]), float(otherPtrUnion.B1[idx]))); break;
|
||||
case ncclFloat8e5m2: ptrUnion.B1[idx] = rccl_bfloat8(ReduceOp(op, float(ptrUnion.B1[idx]), float(otherPtrUnion.B1[idx]))); break;
|
||||
case ncclBfloat16: ptrUnion.B2[idx] = ReduceOp(op, ptrUnion.B2[idx], otherPtrUnion.B2[idx]); break;
|
||||
default:
|
||||
printf("Unsupported datatype (%s)\n", DataTypeToName(dataType).c_str());
|
||||
@@ -365,11 +365,11 @@ void DivideByInt(PtrUnion& ptrUnion, ncclDataType_t const dataType, size_t const
|
||||
case ncclUint32: ptrUnion.U4[idx] /= divisor; break;
|
||||
case ncclInt64: ptrUnion.I8[idx] /= divisor; break;
|
||||
case ncclUint64: ptrUnion.U8[idx] /= divisor; break;
|
||||
case ncclFp8E4M3: ptrUnion.F1[idx] = (rccl_float8((float)(ptrUnion.F1[idx]) / divisor)); break;
|
||||
case ncclFloat8e4m3: ptrUnion.F1[idx] = (rccl_float8((float)(ptrUnion.F1[idx]) / divisor)); break;
|
||||
case ncclFloat16: ptrUnion.F2[idx] = __float2half(__half2float(ptrUnion.F2[idx])/divisor); break;
|
||||
case ncclFloat32: ptrUnion.F4[idx] /= divisor; break;
|
||||
case ncclFloat64: ptrUnion.F8[idx] /= divisor; break;
|
||||
case ncclFp8E5M2: ptrUnion.B1[idx] = (rccl_bfloat8((float)(ptrUnion.B1[idx]) / divisor)); break;
|
||||
case ncclFloat8e5m2: ptrUnion.B1[idx] = (rccl_bfloat8((float)(ptrUnion.B1[idx]) / divisor)); break;
|
||||
case ncclBfloat16: ptrUnion.B2[idx] = (hip_bfloat16((float)(ptrUnion.B2[idx]) / divisor)); break;
|
||||
default:
|
||||
printf("Unsupported datatype (%s)\n", DataTypeToName(dataType).c_str());
|
||||
|
||||
@@ -121,8 +121,8 @@ typedef enum { ncclInt8 = 0, ncclChar = 0,
|
||||
ncclFloat32 = 7, ncclFloat = 7,
|
||||
ncclFloat64 = 8, ncclDouble = 8,
|
||||
ncclBfloat16 = 9,
|
||||
ncclFp8E4M3 = 10,
|
||||
ncclFp8E5M2 = 11,
|
||||
ncclFloat8e4m3 = 10,
|
||||
ncclFloat8e5m2 = 11,
|
||||
ncclNumTypes = 12 } ncclDataType_t;
|
||||
|
||||
/*
|
||||
|
||||
@@ -389,13 +389,10 @@ 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
|
||||
ncclFloat8e4m3 = 10,
|
||||
ncclFloat8e5m2 = 11,
|
||||
ncclNumTypes = 12
|
||||
} ncclDataType_t;
|
||||
/*! @} */
|
||||
|
||||
/*! @defgroup rccl_api_custom_redop Custom Reduction Operator
|
||||
|
||||
Reference in New Issue
Block a user