From ecd812b2d867c176d75217263c3f30132471e6ad Mon Sep 17 00:00:00 2001 From: Jatin Chaudhary Date: Wed, 15 May 2024 01:42:31 +0100 Subject: [PATCH] SWDEV-460834 - add unsafe atomic add for fp16 and bf16 Change-Id: I6de5c2c425c9f8ac7f6c4e5c83c8b8b7ac8fe4cb --- hipamd/include/hip/amd_detail/amd_hip_bf16.h | 35 +++++++++++++++++++ hipamd/include/hip/amd_detail/amd_hip_fp16.h | 36 ++++++++++++++++++++ 2 files changed, 71 insertions(+) diff --git a/hipamd/include/hip/amd_detail/amd_hip_bf16.h b/hipamd/include/hip/amd_detail/amd_hip_bf16.h index cfaa5412a3..7a0dc39911 100644 --- a/hipamd/include/hip/amd_detail/amd_hip_bf16.h +++ b/hipamd/include/hip/amd_detail/amd_hip_bf16.h @@ -102,6 +102,9 @@ #include "amd_hip_vector_types.h" // float2 etc #include "device_library_decls.h" // ocml conversion functions +#if defined(__clang__) && defined(__HIP__) +#include "amd_hip_atomic.h" +#endif // defined(__clang__) && defined(__HIP__) #include "math_fwd.h" // ocml device functions #define __BF16_DEVICE__ __device__ @@ -1811,4 +1814,36 @@ __BF16_DEVICE_STATIC__ __hip_bfloat162 h2trunc(const __hip_bfloat162 h) { __hip_bfloat162_raw hr = h; return __hip_bfloat162(htrunc(__hip_bfloat16_raw{hr.x}), htrunc(__hip_bfloat16_raw{hr.y})); } + +#if defined(__clang__) && defined(__HIP__) +/** + * \ingroup HIP_INTRINSIC_BFLOAT162_MATH + * \brief Atomic add bfloat162 + */ +__BF16_DEVICE_STATIC__ __hip_bfloat162 unsafeAtomicAdd(__hip_bfloat162* address, + __hip_bfloat162 value) { +#if defined(__AMDGCN_UNSAFE_FP_ATOMICS__) && __has_builtin(__builtin_amdgcn_flat_atomic_fadd_v2bf16) + typedef short __attribute__((ext_vector_type(2))) vec_short2; + __hip_bfloat162_raw bf2_v = value; + vec_short2 s2_in{bf2_v.x, bf2_v.y}; + vec_short2 s2_ret = __builtin_amdgcn_flat_atomic_fadd_v2bf16((vec_short2*)address, s2_in); + return __hip_bfloat162_raw{s2_ret[0], s2_ret[1]}; +#else + static_assert(sizeof(unsigned int) == sizeof(__hip_bfloat162_raw)); + union u_hold { + __hip_bfloat162_raw h2r; + unsigned int u32; + }; + u_hold old_val, new_val; + old_val.u32 = + __hip_atomic_load((unsigned int*)address, __ATOMIC_RELAXED, __HIP_MEMORY_SCOPE_AGENT); + do { + new_val.h2r = __hadd2(old_val.h2r, value); + } while (!__hip_atomic_compare_exchange_strong((unsigned int*)address, &old_val.u32, new_val.u32, + __ATOMIC_RELAXED, __ATOMIC_RELAXED, + __HIP_MEMORY_SCOPE_AGENT)); + return old_val.h2r; +#endif +} +#endif // defined(__clang__) && defined(__HIP__) #endif diff --git a/hipamd/include/hip/amd_detail/amd_hip_fp16.h b/hipamd/include/hip/amd_detail/amd_hip_fp16.h index 81883afc96..39a522a503 100644 --- a/hipamd/include/hip/amd_detail/amd_hip_fp16.h +++ b/hipamd/include/hip/amd_detail/amd_hip_fp16.h @@ -30,6 +30,9 @@ THE SOFTWARE. #define __HOST_DEVICE__ __host__ __device__ #include #include "hip/amd_detail/host_defines.h" +#if defined(__clang__) && defined(__HIP__) + #include "hip/amd_detail/amd_hip_atomic.h" +#endif // defined(__clang__) && defined(__HIP__) #include #if defined(__cplusplus) #include @@ -1503,6 +1506,39 @@ THE SOFTWARE. static_cast<__half2_raw>(y).data}; } + // Atomic + #if defined(__clang__) && defined(__HIP__) + inline __device__ __half2 unsafeAtomicAdd(__half2* address, __half2 value) { + #if defined(__AMDGCN_UNSAFE_FP_ATOMICS__) && __has_builtin(__builtin_amdgcn_flat_atomic_fadd_v2f16) + // The api expects an ext_vector_type of half + typedef __fp16 __attribute__((ext_vector_type(2))) vec_fp162; + static_assert(sizeof(vec_fp162) == sizeof(__half2_raw)); + union { + __half2_raw h2r; + vec_fp162 fp16; + } u {value}; + vec_fp162 ret = + __builtin_amdgcn_flat_atomic_fadd_v2f16((vec_fp162*)address, u.fp16); + return __half2{ret[0], ret[1]}; + #else + static_assert(sizeof(__half2_raw) == sizeof(unsigned int)); + union u_hold { + __half2_raw h2r; + unsigned int u32; + }; + u_hold old_val, new_val; + old_val.u32 = __hip_atomic_load((unsigned int*)address, __ATOMIC_RELAXED, + __HIP_MEMORY_SCOPE_AGENT); + do { + new_val.h2r = __hadd2(old_val.h2r, value); + } while (!__hip_atomic_compare_exchange_strong( + (unsigned int*)address, &old_val.u32, new_val.u32, __ATOMIC_RELAXED, + __ATOMIC_RELAXED, __HIP_MEMORY_SCOPE_AGENT)); + return old_val.h2r; + #endif + } + #endif // defined(__clang__) && defined(__HIP__) + // Math functions #if defined(__clang__) && defined(__HIP__) inline