SWDEV-460834 - add unsafe atomic add for fp16 and bf16

Change-Id: I6de5c2c425c9f8ac7f6c4e5c83c8b8b7ac8fe4cb
Этот коммит содержится в:
Jatin Chaudhary
2024-05-15 01:42:31 +01:00
коммит произвёл Jatin Jaikishan Chaudhary
родитель 9d628a4a3d
Коммит ecd812b2d8
2 изменённых файлов: 71 добавлений и 0 удалений
+35
Просмотреть файл
@@ -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
+36
Просмотреть файл
@@ -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