replacing rccl_bfloat16 with hip_bfloat16 (#1126)
Co-authored-by: mberenjk <mberenjk@amd.com>
Этот коммит содержится в:
@@ -167,7 +167,7 @@ namespace RcclUnitTesting
|
||||
case ncclFloat32: F4[idx] = valueF; break;
|
||||
case ncclFloat64: F8[idx] = valueF; break;
|
||||
case ncclFp8E5M2: B1[idx] = rccl_bfloat8(valueF); break;
|
||||
case ncclBfloat16: B2[idx] = rccl_bfloat16(static_cast<float>(valueF)); break;
|
||||
case ncclBfloat16: B2[idx] = hip_bfloat16(static_cast<float>(valueF)); break;
|
||||
default:
|
||||
ERROR("Unsupported datatype\n");
|
||||
return TEST_FAIL;
|
||||
@@ -286,7 +286,7 @@ namespace RcclUnitTesting
|
||||
case ncclFloat32: F4[idx] /= divisor; break;
|
||||
case ncclFloat64: F8[idx] /= divisor; break;
|
||||
case ncclFp8E5M2: B1[idx] = (rccl_bfloat8((float)(B1[idx]) / divisor)); break;
|
||||
case ncclBfloat16: B2[idx] = (rccl_bfloat16((float)(B2[idx]) / divisor)); break;
|
||||
case ncclBfloat16: B2[idx] = (hip_bfloat16((float)(B2[idx]) / divisor)); break;
|
||||
default:
|
||||
ERROR("Unsupported datatype\n");
|
||||
return TEST_FAIL;
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
#include "ErrCode.hpp"
|
||||
#include "rccl/rccl.h"
|
||||
#include "rccl_float8.h"
|
||||
#include "rccl_bfloat16.h"
|
||||
#include <hip/hip_bfloat16.h>
|
||||
#include "hip/hip_fp16.h"
|
||||
|
||||
namespace RcclUnitTesting
|
||||
@@ -48,7 +48,7 @@ namespace RcclUnitTesting
|
||||
float* F4; // ncclFloat32
|
||||
double* F8; // ncclFloat64
|
||||
rccl_bfloat8* B1; // ncclFp8E5M2
|
||||
rccl_bfloat16* B2; // ncclBfloat16
|
||||
hip_bfloat16* B2; // ncclBfloat16
|
||||
|
||||
constexpr PtrUnion() : ptr(nullptr) {}
|
||||
|
||||
|
||||
Ссылка в новой задаче
Block a user