Merge remote-tracking branch 'nccl/master' into develop

This commit is contained in:
BertanDogancay
2025-04-23 20:46:36 -07:00
118 ha cambiato i file con 12789 aggiunte e 3358 eliminazioni
+12 -12
Vedi File
@@ -195,18 +195,18 @@ namespace RcclUnitTesting
scalarsPerRank.Attach(scalarsPerRank.ptr);
switch (this->dataType)
{
case ncclInt8: ss << scalarsPerRank.I1[this->globalRank]; break;
case ncclUint8: ss << scalarsPerRank.U1[this->globalRank]; break;
case ncclInt32: ss << scalarsPerRank.I4[this->globalRank]; break;
case ncclUint32: ss << scalarsPerRank.U4[this->globalRank]; break;
case ncclInt64: ss << scalarsPerRank.I8[this->globalRank]; break;
case ncclUint64: ss << scalarsPerRank.U8[this->globalRank]; break;
case ncclFp8E4M3: ss << scalarsPerRank.F1[this->globalRank]; break;
case ncclFloat32: ss << scalarsPerRank.F4[this->globalRank]; break;
case ncclFloat64: ss << scalarsPerRank.F8[this->globalRank]; break;
case ncclFp8E5M2: ss << scalarsPerRank.B1[this->globalRank]; break;
case ncclBfloat16: ss << scalarsPerRank.B2[this->globalRank]; break;
default: ss << "(UNKNOWN)";
case ncclInt8: ss << scalarsPerRank.I1[this->globalRank]; break;
case ncclUint8: ss << scalarsPerRank.U1[this->globalRank]; break;
case ncclInt32: ss << scalarsPerRank.I4[this->globalRank]; break;
case ncclUint32: ss << scalarsPerRank.U4[this->globalRank]; break;
case ncclInt64: ss << scalarsPerRank.I8[this->globalRank]; break;
case ncclUint64: ss << scalarsPerRank.U8[this->globalRank]; break;
case ncclFloat8e4m3: ss << scalarsPerRank.F1[this->globalRank]; break;
case ncclFloat32: ss << scalarsPerRank.F4[this->globalRank]; break;
case ncclFloat64: ss << scalarsPerRank.F8[this->globalRank]; break;
case ncclFloat8e5m2: ss << scalarsPerRank.B1[this->globalRank]; break;
case ncclBfloat16: ss << scalarsPerRank.B2[this->globalRank]; break;
default: ss << "(UNKNOWN)";
}
ss << " ";
}
+2 -2
Vedi File
@@ -54,8 +54,8 @@ namespace RcclUnitTesting
"ncclFloat32",
"ncclFloat64",
"ncclBfloat16",
"ncclFp8E4M3",
"ncclFp8E5M2"
"ncclFloat8e4m3",
"ncclFloat8e5m2"
};
char const ncclRedOpNames[ncclNumOps][32] =
+2 -2
Vedi File
@@ -282,8 +282,8 @@ namespace RcclUnitTesting
dataTypes.push_back(ncclFloat32);
dataTypes.push_back(ncclFloat64);
dataTypes.push_back(ncclBfloat16);
dataTypes.push_back(ncclFp8E4M3);
dataTypes.push_back(ncclFp8E5M2);
dataTypes.push_back(ncclFloat8e4m3);
dataTypes.push_back(ncclFloat8e5m2);
}
// Build list of possible # GPU ranks based on env vars
+21 -21
Vedi File
@@ -14,8 +14,8 @@ namespace RcclUnitTesting
{
case ncclInt8: return 1;
case ncclUint8: return 1;
case ncclFp8E4M3:return 1;
case ncclFp8E5M2:return 1;
case ncclFloat8e4m3:return 1;
case ncclFloat8e5m2:return 1;
case ncclInt32: return 4;
case ncclUint32: return 4;
case ncclInt64: return 8;
@@ -150,7 +150,7 @@ namespace RcclUnitTesting
{
// Due to floating-point math not being commutative, the ordering in which ranks are added will matter.
// For lower-precision data types, we initialize all ranks to the same value to avoid this
int valueI = (dataType == ncclFp8E4M3 || dataType == ncclFp8E5M2)? (i % 16) :(globalRank + i) % 256;
int valueI = (dataType == ncclFloat8e4m3 || dataType == ncclFloat8e5m2)? (i % 16) :(globalRank + i) % 256;
double valueF = 1.0L/((double)valueI+1.0L);
temp.Set(dataType, i, valueI, valueF);
}
@@ -179,11 +179,11 @@ namespace RcclUnitTesting
case ncclUint32: U4[idx] = valueI; break;
case ncclInt64: I8[idx] = valueI; break;
case ncclUint64: U8[idx] = valueI; break;
case ncclFp8E4M3: F1[idx] = rccl_float8(valueF); break;
case ncclFloat8e4m3: F1[idx] = rccl_float8(valueF); break;
case ncclFloat16: F2[idx] = __float2half(static_cast<float>(valueF)); break;
case ncclFloat32: F4[idx] = valueF; break;
case ncclFloat64: F8[idx] = valueF; break;
case ncclFp8E5M2: B1[idx] = rccl_bfloat8(valueF); break;
case ncclFloat8e5m2: B1[idx] = rccl_bfloat8(valueF); break;
case ncclBfloat16: B2[idx] = hip_bfloat16(static_cast<float>(valueF)); break;
default:
ERROR("Unsupported datatype\n");
@@ -202,11 +202,11 @@ namespace RcclUnitTesting
case ncclUint32: valueI = U4[idx]; break;
case ncclInt64: valueI = I8[idx]; break;
case ncclUint64: valueI = U8[idx]; break;
case ncclFp8E4M3: valueF = float(F1[idx]); break;
case ncclFloat8e4m3: valueF = float(F1[idx]); break;
case ncclFloat16: valueF = __half2float(F2[idx]); break;
case ncclFloat32: valueF = F4[idx]; break;
case ncclFloat64: valueF = F8[idx]; break;
case ncclFp8E5M2: valueF = float(B1[idx]); break;
case ncclFloat8e5m2: valueF = float(B1[idx]); break;
case ncclBfloat16: valueF = B2[idx]; break;
default:
ERROR("Unsupported datatype\n");
@@ -234,11 +234,11 @@ namespace RcclUnitTesting
case ncclUint32: U4[idx] *= scalarsPerRank.U4[rank]; break;
case ncclInt64: I8[idx] *= scalarsPerRank.I8[rank]; break;
case ncclUint64: U8[idx] *= scalarsPerRank.U8[rank]; break;
case ncclFp8E4M3: F1[idx] = rccl_float8(F1[idx] * scalarsPerRank.F1[rank]); break;
case ncclFloat8e4m3: F1[idx] = rccl_float8(F1[idx] * scalarsPerRank.F1[rank]); break;
case ncclFloat16: F2[idx] = __float2half(__half2float(F2[idx]) * __half2float(scalarsPerRank.F2[rank])); break;
case ncclFloat32: F4[idx] *= scalarsPerRank.F4[rank]; break;
case ncclFloat64: F8[idx] *= scalarsPerRank.F8[rank]; break;
case ncclFp8E5M2: B1[idx] = rccl_bfloat8(B1[idx] * scalarsPerRank.B1[rank]); break;
case ncclFloat8e5m2: B1[idx] = rccl_bfloat8(B1[idx] * scalarsPerRank.B1[rank]); break;
case ncclBfloat16: B2[idx] *= scalarsPerRank.B2[rank]; break;
default:
ERROR("Unsupported datatype\n");
@@ -269,11 +269,11 @@ namespace RcclUnitTesting
case ncclUint32: U4[idx] = ReduceOp(op, U4[idx], inputCpu.U4[idx]); break;
case ncclInt64: I8[idx] = ReduceOp(op, I8[idx], inputCpu.I8[idx]); break;
case ncclUint64: U8[idx] = ReduceOp(op, U8[idx], inputCpu.U8[idx]); break;
case ncclFp8E4M3: F1[idx] = rccl_float8(ReduceOp(op, float(F1[idx]), float(inputCpu.F1[idx]))); break;
case ncclFloat8e4m3: F1[idx] = rccl_float8(ReduceOp(op, float(F1[idx]), float(inputCpu.F1[idx]))); break;
case ncclFloat16: F2[idx] = __float2half(ReduceOp(op, __half2float(F2[idx]), __half2float(inputCpu.F2[idx]))); break;
case ncclFloat32: F4[idx] = ReduceOp(op, F4[idx], inputCpu.F4[idx]); break;
case ncclFloat64: F8[idx] = ReduceOp(op, F8[idx], inputCpu.F8[idx]); break;
case ncclFp8E5M2: B1[idx] = rccl_bfloat8(ReduceOp(op, float(B1[idx]), float(inputCpu.B1[idx]))); break;
case ncclFloat8e5m2: B1[idx] = rccl_bfloat8(ReduceOp(op, float(B1[idx]), float(inputCpu.B1[idx]))); break;
case ncclBfloat16: B2[idx] = ReduceOp(op, B2[idx], inputCpu.B2[idx]); break;
default:
ERROR("Unsupported datatype\n");
@@ -298,11 +298,11 @@ namespace RcclUnitTesting
case ncclUint32: U4[idx] /= divisor; break;
case ncclInt64: I8[idx] /= divisor; break;
case ncclUint64: U8[idx] /= divisor; break;
case ncclFp8E4M3: F1[idx] = (rccl_float8((float)(F1[idx]) / divisor)); break;
case ncclFloat8e4m3: F1[idx] = (rccl_float8((float)(F1[idx]) / divisor)); break;
case ncclFloat16: F2[idx] = __float2half(__half2float(F2[idx])/divisor); break;
case ncclFloat32: F4[idx] /= divisor; break;
case ncclFloat64: F8[idx] /= divisor; break;
case ncclFp8E5M2: B1[idx] = (rccl_bfloat8((float)(B1[idx]) / divisor)); break;
case ncclFloat8e5m2: B1[idx] = (rccl_bfloat8((float)(B1[idx]) / divisor)); break;
case ncclBfloat16: B2[idx] = (hip_bfloat16((float)(B2[idx]) / divisor)); break;
default:
ERROR("Unsupported datatype\n");
@@ -330,11 +330,11 @@ namespace RcclUnitTesting
case ncclUint32: isMatch = (U4[idx] == expected.U4[idx]); break;
case ncclInt64: isMatch = (I8[idx] == expected.I8[idx]); break;
case ncclUint64: isMatch = (U8[idx] == expected.U8[idx]); break;
case ncclFp8E4M3: isMatch = (fabs(float(F1[idx]) - float(expected.F1[idx])) < 9e-2); break;
case ncclFloat8e4m3: isMatch = (fabs(float(F1[idx]) - float(expected.F1[idx])) < 9e-2); break;
case ncclFloat16: isMatch = (fabs(__half2float(F2[idx]) - __half2float(expected.F2[idx])) < 9e-2); break;
case ncclFloat32: isMatch = (fabs(F4[idx] - expected.F4[idx]) < 1e-5); break;
case ncclFloat64: isMatch = (fabs(F8[idx] - expected.F8[idx]) < 1e-12); break;
case ncclFp8E5M2: isMatch = (fabs(float(B1[idx]) - float(expected.B1[idx])) < 9e-2); break;
case ncclFloat8e5m2: isMatch = (fabs(float(B1[idx]) - float(expected.B1[idx])) < 9e-2); break;
case ncclBfloat16: isMatch = (fabs((float)B2[idx] - (float)expected.B2[idx]) < 9e-2); break;
default:
ERROR("Unsupported datatype\n");
@@ -359,16 +359,16 @@ namespace RcclUnitTesting
ERROR("Expected output: %ld. Actual output: %ld at index %lu\n", expected.I8[idx], I8[idx], idx); break;
case ncclUint64:
ERROR("Expected output: %lu. Actual output: %lu at index %lu\n", expected.U8[idx], U8[idx], idx); break;
case ncclFp8E4M3:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", (float)expected.F1[idx], (float)F1[idx], idx); break;
case ncclFloat8e4m3:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", (float)expected.F1[idx], (float)F1[idx], idx);
case ncclFloat16:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", __half2float(expected.F2[idx]), __half2float(F2[idx]), idx); break;
case ncclFloat32:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", expected.F4[idx], F4[idx], idx); break;
case ncclFloat64:
ERROR("Expected output: %lf. Actual output: %lf at index %lu\n", expected.F8[idx], F8[idx], idx); break;
case ncclFp8E5M2:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", (float)expected.B1[idx], (float)B1[idx], idx); break;
case ncclFloat8e5m2:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", (float)expected.B1[idx], (float)B1[idx], idx);
case ncclBfloat16:
ERROR("Expected output: %f. Actual output: %f at index %lu\n", (float)expected.B2[idx], (float)B2[idx], idx); break;
default:
@@ -393,11 +393,11 @@ namespace RcclUnitTesting
case ncclUint32: ss << U4[i]; break;
case ncclInt64: ss << I8[i]; break;
case ncclUint64: ss << U8[i]; break;
case ncclFp8E4M3: ss << (float)F1[i]; break;
case ncclFloat8e4m3: ss << (float)F1[i]; break;
case ncclFloat16: ss << __half2float(F2[i]); break;
case ncclFloat32: ss << F4[i]; break;
case ncclFloat64: ss << F8[i]; break;
case ncclFp8E5M2: ss << (float)B1[i]; break;
case ncclFloat8e5m2: ss << (float)B1[i]; break;
case ncclBfloat16: ss << (float)B2[i]; break;
default: break;
}
+2 -2
Vedi File
@@ -44,10 +44,10 @@ namespace RcclUnitTesting
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) {}
+2 -2
Vedi File
@@ -697,8 +697,8 @@ namespace RcclUnitTesting
{
//Skipping AllReduce FP8 test on 9 to 16 ranks (gfx90a).
if(ev.isGfx90 && numRanks > 8 && funcTypes[ftIdx] == ncclCollAllReduce
&& (dataTypes[dtIdx] == ncclFp8E4M3
|| dataTypes[dtIdx] == ncclFp8E5M2))
&& (dataTypes[dtIdx] == ncclFloat8e4m3
|| dataTypes[dtIdx] == ncclFloat8e5m2))
{
continue;
}