SWDEV-460834 - add unsafe atomic add for fp16 and bf16
Change-Id: I6de5c2c425c9f8ac7f6c4e5c83c8b8b7ac8fe4cb
Этот коммит содержится в:
коммит произвёл
Jatin Jaikishan Chaudhary
родитель
9d628a4a3d
Коммит
ecd812b2d8
@@ -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
|
||||
|
||||
@@ -30,6 +30,9 @@ THE SOFTWARE.
|
||||
#define __HOST_DEVICE__ __host__ __device__
|
||||
#include <hip/amd_detail/amd_hip_common.h>
|
||||
#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 <assert.h>
|
||||
#if defined(__cplusplus)
|
||||
#include <algorithm>
|
||||
@@ -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
|
||||
|
||||
Ссылка в новой задаче
Block a user