From 0580e2053ce2899b578e3ba3d6d7310260d0f0f4 Mon Sep 17 00:00:00 2001 From: AidanBeltonS Date: Mon, 24 Nov 2025 09:14:03 +0000 Subject: [PATCH] SWDEV-533546, SWDEV-540027 - Add e8m0 conversions and testing (#987) * SWDEV-533546 - Add conversion functions for e8m0 * SWDEV-533546 - remove whitespace * Add testing * Update based on feedback * Copilot suggestions --------- Co-authored-by: systems-assistant[bot] --- .../include/hip/amd_detail/amd_hip_fp8.h | 117 ++++++ .../catch/unit/deviceLib/CMakeLists.txt | 1 + .../catch/unit/deviceLib/fp8_e8m0.cc | 344 ++++++++++++++++++ 3 files changed, 462 insertions(+) create mode 100644 projects/hip-tests/catch/unit/deviceLib/fp8_e8m0.cc diff --git a/projects/clr/hipamd/include/hip/amd_detail/amd_hip_fp8.h b/projects/clr/hipamd/include/hip/amd_detail/amd_hip_fp8.h index 13d5cb92a2..8e4487d92d 100644 --- a/projects/clr/hipamd/include/hip/amd_detail/amd_hip_fp8.h +++ b/projects/clr/hipamd/include/hip/amd_detail/amd_hip_fp8.h @@ -65,6 +65,7 @@ // Include it explicitly for HIPRTC #include "amd_hip_bf16.h" +#include "amd_hip_mx_common.h" #if !defined(__HIPCC_RTC__) #include @@ -950,6 +951,122 @@ __FP8_HOST_STATIC__ __hip_fp8x2_storage_t __hip_cvt_halfraw2_to_fp8x2( return __hip_cvt_float2_to_fp8x2(__half22float2(__half2(x)), sat, interp); } +namespace hip_detail { + +constexpr __hip_fp8_storage_t e8m0_NaN = 0xFFU; +constexpr __hip_internal::uint16_t bf16_NaN = 0x7FFFU; + +constexpr __hip_internal::uint16_t bf16_sig_mask = 0x007FU; +constexpr __hip_internal::uint32_t float_sig_mask = 0x007FFFFFU; +constexpr __hip_internal::uint64_t double_sig_mask = 0x000FFFFFFFFFFFFFU; + +constexpr __hip_internal::uint16_t bf16_max_exp = 0x7F80U; +constexpr __hip_internal::uint32_t float_max_exp = 0x7F800000U; +constexpr __hip_internal::uint64_t double_max_exp = 0x7FF0000000000000U; + +constexpr __hip_internal::uint16_t bf16_sign_mask = 0x8000U; +constexpr __hip_internal::uint32_t float_sign_mask = 0x80000000U; +constexpr __hip_internal::uint64_t double_sign_mask = 0x8000000000000000U; + +constexpr __hip_internal::uint16_t bf16_half_sig_bit = 0x0040U; +constexpr __hip_internal::uint32_t float_half_sig_bit = 0x00400000U; +constexpr __hip_internal::uint64_t double_half_sig_bit = 0x0008000000000000U; + +} // namespace hip_detail + +__FP8_HOST_DEVICE_STATIC__ __hip_fp8_storage_t __hip_cvt_double_to_e8m0( + const double val, const __hip_saturation_t saturate, const enum hipRoundMode rounding) { + union { + double as_double; + __hip_internal::uint64_t as_int; + } u{val}; + + // Shifts out mantissa bits from double dtype + unsigned short double_exp = + static_cast((~hip_detail::double_sign_mask & u.as_int) >> 52); + __hip_fp8_storage_t e8m0; + if (double_exp == 0x0U) { + e8m0 = 0x0U; + } else if ((double_exp - 0x0380U) > 0x00FF) { // if double is NaN/Inf/or too large + e8m0 = hip_detail::e8m0_NaN; + } else { + e8m0 = + double_exp - 0x0380U; // shift due to bias difference between double and single precision + } + + // If there is a mantissa and the exp wont overflow round up + if ((rounding == hipRoundPosInf) && (u.as_int & hip_detail::double_sig_mask) && + (!((u.as_int & ~hip_detail::double_sign_mask) < hip_detail::double_half_sig_bit)) && + (e8m0 < hip_detail::e8m0_NaN)) { + ++e8m0; + } + + // If e8m0 is NaN and exponent is a large non-inf value round down to a value + if ((saturate == __HIP_SATFINITE) && (e8m0 == hip_detail::e8m0_NaN) && + ((u.as_int & ~hip_detail::double_sign_mask) <= hip_detail::double_max_exp)) { + --e8m0; + } + return e8m0; +} + +__FP8_HOST_DEVICE_STATIC__ __hip_fp8_storage_t __hip_cvt_float_to_e8m0( + const float val, const __hip_saturation_t saturate, const enum hipRoundMode rounding) { + union { + float as_float; + __hip_internal::uint32_t as_int; + } u{val}; + + // Shifts out mantissa bits from float dtype + __hip_fp8_storage_t e8m0 = static_cast(u.as_int >> 23); + + // If there is a mantissa and the exp wont overflow round up + if ((rounding == hipRoundPosInf) && (u.as_int & hip_detail::float_sig_mask) && + (!((u.as_int & ~hip_detail::float_sign_mask) < hip_detail::float_half_sig_bit)) && + (e8m0 < hip_detail::e8m0_NaN)) { + ++e8m0; + } + + // If e8m0 is NaN and exponent is a large non-inf value round down to a value + if ((saturate == __HIP_SATFINITE) && (e8m0 == hip_detail::e8m0_NaN) && + ((u.as_int & ~hip_detail::float_sign_mask) <= hip_detail::float_max_exp)) { + --e8m0; + } + return e8m0; +} + +__FP8_HOST_DEVICE_STATIC__ __hip_fp8_storage_t +__hip_cvt_bfloat16raw_to_e8m0(const __hip_bfloat16_raw hr, const __hip_saturation_t saturate, + const enum hipRoundMode rounding) { + // Shifts out mantissa bits from bf16 dtype + __hip_fp8_storage_t e8m0 = static_cast(hr.x >> 7); + + // If there is a mantissa and the exp wont overflow round up + if ((rounding == hipRoundPosInf) && (hr.x & hip_detail::bf16_sig_mask) && + (!((hr.x & ~hip_detail::bf16_sign_mask) < hip_detail::bf16_half_sig_bit)) && + (e8m0 < hip_detail::e8m0_NaN)) { + ++e8m0; + } + + // If e8m0 is NaN and exponent is a large non-inf value round down to a value + if ((saturate == __HIP_SATFINITE) && (e8m0 == hip_detail::e8m0_NaN) && + ((hr.x & ~hip_detail::bf16_sign_mask) <= hip_detail::bf16_max_exp)) { + --e8m0; + } + return e8m0; +} + +__FP8_HOST_DEVICE_STATIC__ __hip_bfloat16_raw +__hip_cvt_e8m0_to_bf16raw(const __hip_fp8_storage_t x) { + switch (x) { + case 0x00U: + return __hip_bfloat16_raw{0x0040U}; + case hip_detail::e8m0_NaN: + return __hip_bfloat16_raw{hip_detail::bf16_NaN}; + default: + return __hip_bfloat16_raw{static_cast(x << 7)}; + } +} + /** * \brief struct representing single fp8 number with e4m3 interpretation * diff --git a/projects/hip-tests/catch/unit/deviceLib/CMakeLists.txt b/projects/hip-tests/catch/unit/deviceLib/CMakeLists.txt index 1ad4dba46b..9b39b0a521 100644 --- a/projects/hip-tests/catch/unit/deviceLib/CMakeLists.txt +++ b/projects/hip-tests/catch/unit/deviceLib/CMakeLists.txt @@ -85,6 +85,7 @@ set(AMD_TEST_SRC AtomicsWithRandomActiveLanesInWavefront.cc fp16_ops.cc fp8_host.cc + fp8_e8m0.cc fp6_ocp.cc fp4_ocp.cc ) diff --git a/projects/hip-tests/catch/unit/deviceLib/fp8_e8m0.cc b/projects/hip-tests/catch/unit/deviceLib/fp8_e8m0.cc new file mode 100644 index 0000000000..e094ed2bb2 --- /dev/null +++ b/projects/hip-tests/catch/unit/deviceLib/fp8_e8m0.cc @@ -0,0 +1,344 @@ +/* +Copyright (c) 2025 Advanced Micro Devices, Inc. All rights reserved. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +THE SOFTWARE. +*/ + +#include +#include +#include + +void host_cvt_bfloat16raw_to_e8m0(const std::vector<__hip_bfloat16>& in, + std::vector& out, __hip_saturation_t sat, + hipRoundMode round) { + for (size_t i = 0; i < in.size(); ++i) { + out[i] = __hip_cvt_bfloat16raw_to_e8m0(in[i], sat, round); + } +} + +__global__ void bfloat16raw_to_e8m0(const __hip_bfloat16* in, unsigned char* out, size_t size, + __hip_saturation_t sat, hipRoundMode round) { + size_t tid = threadIdx.x + blockDim.x * blockIdx.x; + + if (tid < size) { + out[tid] = __hip_cvt_bfloat16raw_to_e8m0(in[tid], sat, round); + } +} +void device_cvt_bfloat16raw_to_e8m0(const std::vector<__hip_bfloat16>& in, + std::vector& out, __hip_saturation_t sat, + hipRoundMode round) { + __hip_bfloat16* in_d = nullptr; + unsigned char* out_d = nullptr; + REQUIRE(in.size() < 1024); + + HIP_CHECK(hipMalloc(&in_d, sizeof(__hip_bfloat16) * in.size())); + HIP_CHECK(hipMalloc(&out_d, sizeof(unsigned char) * out.size())); + + HIP_CHECK(hipMemcpy(in_d, in.data(), sizeof(__hip_bfloat16) * in.size(), hipMemcpyHostToDevice)); + + bfloat16raw_to_e8m0<<<1, 1024>>>(in_d, out_d, in.size(), sat, round); + + HIP_CHECK( + hipMemcpy(out.data(), out_d, sizeof(unsigned char) * out.size(), hipMemcpyDeviceToHost)); +} + +TEST_CASE("Unit__hip_cvt_bfloat16raw_to_e8m0") { + bool run_on_host = GENERATE(true, false); + __hip_saturation_t saturation = GENERATE(__HIP_NOSAT, __HIP_SATFINITE); + hipRoundMode rounding = GENERATE(hipRoundZero, hipRoundPosInf); + + std::vector<__hip_bfloat16> in = {0.0f, + 0.5f, + 0.6f, + 4, + 5, + 8, + -0.5f, + -0.6f, + -4, + -5, + -8, + 1e38, + -1e38, + std::nanf("1"), + -std::nanf("1"), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity()}; + std::vector exp(in.size()); + std::vector out(exp.size()); + + if (rounding == hipRoundPosInf) { + exp = {0x00U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0xFE, 0xFE}; + } else { + exp = {0x00U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0xFD, 0xFD}; + } + if (saturation == __HIP_NOSAT) { + std::vector exp_append = {0xFF, 0xFF, 0xFF, 0xFF}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } else { + std::vector exp_append = {0xFF, 0xFF, 0xFE, 0xFE}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } + + REQUIRE(exp.size() == in.size()); + + if (run_on_host) { + host_cvt_bfloat16raw_to_e8m0(in, out, saturation, rounding); + } else { + device_cvt_bfloat16raw_to_e8m0(in, out, saturation, rounding); + } + + for (size_t i = 0; i < in.size(); ++i) { + INFO("out:" << out[i] << " exp:" << exp[i] << " for index:" << i); + REQUIRE(out[i] == exp[i]); + } +} + +//// + +void host_cvt_float_to_e8m0(const std::vector& in, std::vector& out, + __hip_saturation_t sat, hipRoundMode round) { + for (size_t i = 0; i < in.size(); ++i) { + out[i] = __hip_cvt_float_to_e8m0(in[i], sat, round); + } +} + +__global__ void float_to_e8m0_kernel(const float* in, unsigned char* out, size_t size, + __hip_saturation_t sat, hipRoundMode round) { + size_t tid = threadIdx.x + blockDim.x * blockIdx.x; + + if (tid < size) { + out[tid] = __hip_cvt_float_to_e8m0(in[tid], sat, round); + } +} +void device_cvt_float_to_e8m0(const std::vector& in, std::vector& out, + __hip_saturation_t sat, hipRoundMode round) { + float* in_d = nullptr; + unsigned char* out_d = nullptr; + REQUIRE(in.size() < 1024); + + HIP_CHECK(hipMalloc(&in_d, sizeof(float) * in.size())); + HIP_CHECK(hipMalloc(&out_d, sizeof(unsigned char) * out.size())); + + HIP_CHECK(hipMemcpy(in_d, in.data(), sizeof(float) * in.size(), hipMemcpyHostToDevice)); + + float_to_e8m0_kernel<<<1, 1024>>>(in_d, out_d, in.size(), sat, round); + + HIP_CHECK( + hipMemcpy(out.data(), out_d, sizeof(unsigned char) * out.size(), hipMemcpyDeviceToHost)); +} + +TEST_CASE("Unit__hip_cvt_float_to_e8m0") { + bool run_on_host = GENERATE(true, false); + __hip_saturation_t saturation = GENERATE(__HIP_NOSAT, __HIP_SATFINITE); + hipRoundMode rounding = GENERATE(hipRoundZero, hipRoundPosInf); + + std::vector in = {0.0f, + 0.5f, + 0.6f, + 4.0f, + 5.0f, + 8.0f, + -0.5f, + -0.6f, + -4.0f, + -5.0f, + -8.0f, + 1e38, + -1e38, + std::nanf("1"), + -std::nanf("1"), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity()}; + std::vector exp(in.size()); + std::vector out(exp.size()); + + if (rounding == hipRoundPosInf) { + exp = {0x00U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0xFE, 0xFE}; + } else { + exp = {0x00U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0xFD, 0xFD}; + } + if (saturation == __HIP_NOSAT) { + std::vector exp_append = {0xFF, 0xFF, 0xFF, 0xFF}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } else { + std::vector exp_append = {0xFF, 0xFF, 0xFE, 0xFE}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } + + REQUIRE(exp.size() == in.size()); + + if (run_on_host) { + host_cvt_float_to_e8m0(in, out, saturation, rounding); + } else { + device_cvt_float_to_e8m0(in, out, saturation, rounding); + } + + for (size_t i = 0; i < in.size(); ++i) { + INFO("out:" << out[i] << " exp:" << exp[i] << " for index:" << i); + REQUIRE(out[i] == exp[i]); + } +} + +//// + +void host_cvt_double_to_e8m0(const std::vector& in, std::vector& out, + __hip_saturation_t sat, hipRoundMode round) { + for (size_t i = 0; i < in.size(); ++i) { + out[i] = __hip_cvt_double_to_e8m0(in[i], sat, round); + } +} + +__global__ void double_to_e8m0_kernel(const double* in, unsigned char* out, size_t size, + __hip_saturation_t sat, hipRoundMode round) { + size_t tid = threadIdx.x + blockDim.x * blockIdx.x; + + if (tid < size) { + out[tid] = __hip_cvt_double_to_e8m0(in[tid], sat, round); + } +} +void device_cvt_double_to_e8m0(const std::vector& in, std::vector& out, + __hip_saturation_t sat, hipRoundMode round) { + double* in_d = nullptr; + unsigned char* out_d = nullptr; + REQUIRE(in.size() < 1024); + + HIP_CHECK(hipMalloc(&in_d, sizeof(double) * in.size())); + HIP_CHECK(hipMalloc(&out_d, sizeof(unsigned char) * out.size())); + + HIP_CHECK(hipMemcpy(in_d, in.data(), sizeof(double) * in.size(), hipMemcpyHostToDevice)); + + double_to_e8m0_kernel<<<1, 1024>>>(in_d, out_d, in.size(), sat, round); + + HIP_CHECK( + hipMemcpy(out.data(), out_d, sizeof(unsigned char) * out.size(), hipMemcpyDeviceToHost)); +} + + +TEST_CASE("Unit__hip_cvt_double_to_e8m0") { + bool run_on_host = GENERATE(true, false); + __hip_saturation_t saturation = GENERATE(__HIP_NOSAT, __HIP_SATFINITE); + hipRoundMode rounding = GENERATE(hipRoundZero, hipRoundPosInf); + + std::vector in = {0.0, + 0.5, + 0.6, + 4.0, + 5.0, + 8.0, + -0.5, + -0.6, + -4.0, + -5.0, + -8.0, + 1e38, + -1e38, + std::nan("1"), + -std::nan("1"), + std::numeric_limits::infinity(), + -std::numeric_limits::infinity(), + 1e50, + -1e50}; + std::vector exp(in.size()); + std::vector out(exp.size()); + + if (rounding == hipRoundPosInf) { + exp = {0x00U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0x7EU, 0x7FU, 0x81U, 0x82U, 0x82U, 0xFE, 0xFE}; + } else { + exp = {0x00U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0x7EU, 0x7EU, 0x81U, 0x81U, 0x82U, 0xFD, 0xFD}; + } + if (saturation == __HIP_NOSAT) { + std::vector exp_append = {0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } else { + std::vector exp_append = {0xFF, 0xFF, 0xFE, 0xFE, 0xFE, 0xFE}; + exp.insert(exp.end(), exp_append.begin(), exp_append.end()); + } + + REQUIRE(exp.size() == in.size()); + + if (run_on_host) { + host_cvt_double_to_e8m0(in, out, saturation, rounding); + } else { + device_cvt_double_to_e8m0(in, out, saturation, rounding); + } + + for (size_t i = 0; i < in.size(); ++i) { + INFO("out:" << out[i] << " exp:" << exp[i] << " for index:" << i); + REQUIRE(out[i] == exp[i]); + } +} + +//// + +void host_cvt_e8m0_to_bf16raw(const std::vector& in, std::vector& out) { + for (size_t i = 0; i < in.size(); ++i) { + __hip_bfloat16 temp = __hip_cvt_e8m0_to_bf16raw(in[i]); + out[i] = static_cast(temp); + } +} + +__global__ void e8m0_to_bf16raw_kernel(const unsigned char* in, float* out, size_t size) { + size_t tid = threadIdx.x + blockDim.x * blockIdx.x; + + if (tid < size) { + __hip_bfloat16 temp = __hip_cvt_e8m0_to_bf16raw(in[tid]); + out[tid] = static_cast(temp); + } +} +void device_cvt_e8m0_to_bf16raw(const std::vector& in, std::vector& out) { + unsigned char* in_d = nullptr; + float* out_d = nullptr; + REQUIRE(in.size() < 1024); + + HIP_CHECK(hipMalloc(&in_d, sizeof(unsigned char) * in.size())); + HIP_CHECK(hipMalloc(&out_d, sizeof(float) * out.size())); + + HIP_CHECK(hipMemcpy(in_d, in.data(), sizeof(unsigned char) * in.size(), hipMemcpyHostToDevice)); + + e8m0_to_bf16raw_kernel<<<1, 1024>>>(in_d, out_d, in.size()); + + HIP_CHECK(hipMemcpy(out.data(), out_d, sizeof(float) * out.size(), hipMemcpyDeviceToHost)); +} + +TEST_CASE("Unit__hip_cvt_e8m0_to_bf16raw") { + bool run_on_host = GENERATE(true, false); + + std::vector in = {0x00u, 0x7EU, 0x81U, 0x82U, 0xFF}; + std::vector exp = {0.0f, 0.5f, 4.0f, 8.0f, std::nanf("0")}; + std::vector out(exp.size()); + + REQUIRE(exp.size() == in.size()); + + if (run_on_host) { + host_cvt_e8m0_to_bf16raw(in, out); + } else { + device_cvt_e8m0_to_bf16raw(in, out); + } + + for (size_t i = 0; i < in.size(); ++i) { + INFO("out:" << out[i] << " exp:" << exp[i] << " for index:" << i); + if (std::isnan(exp[i])) { + REQUIRE(std::isnan(out[i])); + } else { + REQUIRE_THAT(out[i], Catch::WithinAbs(exp[i], 1e-6f) || Catch::WithinRel(exp[i], 1e-3f)); + } + } +} +