#ifndef OMPTARGET_DEVICERTL_SYNCHRONIZATION_H
#define OMPTARGET_DEVICERTL_SYNCHRONIZATION_H
#include "DeviceTypes.h"
#include "DeviceUtils.h"
namespace ompx {
namespace atomic {
enum OrderingTy {
relaxed = __ATOMIC_RELAXED,
acquire = __ATOMIC_ACQUIRE,
release = __ATOMIC_RELEASE,
acq_rel = __ATOMIC_ACQ_REL,
seq_cst = __ATOMIC_SEQ_CST,
};
enum MemScopeTy {
system = __MEMORY_SCOPE_SYSTEM,
device = __MEMORY_SCOPE_DEVICE,
workgroup = __MEMORY_SCOPE_WRKGRP,
wavefront = __MEMORY_SCOPE_WVFRNT,
single = __MEMORY_SCOPE_SINGLE,
};
uint32_t inc(uint32_t *Addr, uint32_t V, OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device);
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
bool cas(Ty *Address, V ExpectedV, V DesiredV, atomic::OrderingTy OrderingSucc,
atomic::OrderingTy OrderingFail,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_compare_exchange(Address, &ExpectedV, &DesiredV, false,
OrderingSucc, OrderingFail, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V add(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_add(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V load(Ty *Address, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
#ifdef __NVPTX__
return __scoped_atomic_fetch_add(Address, V(0), Ordering, MemScope);
#else
return __scoped_atomic_load_n(Address, Ordering, MemScope);
#endif
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
void store(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
__scoped_atomic_store_n(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V mul(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
Ty TypedCurrentVal, TypedResultVal, TypedNewVal;
bool Success;
do {
TypedCurrentVal = atomic::load(Address, Ordering);
TypedNewVal = TypedCurrentVal * Val;
Success = atomic::cas(Address, TypedCurrentVal, TypedNewVal, Ordering,
atomic::relaxed, MemScope);
} while (!Success);
return TypedResultVal;
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<!utils::is_floating_point_v<V>, V>
max(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_max(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<utils::is_same_v<V, float>, V>
max(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
if (Val >= 0)
return utils::bitCast<float>(max(
(int32_t *)Address, utils::bitCast<int32_t>(Val), Ordering, MemScope));
return utils::bitCast<float>(min(
(uint32_t *)Address, utils::bitCast<uint32_t>(Val), Ordering, MemScope));
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<utils::is_same_v<V, double>, V>
max(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
if (Val >= 0)
return utils::bitCast<double>(max(
(int64_t *)Address, utils::bitCast<int64_t>(Val), Ordering, MemScope));
return utils::bitCast<double>(min(
(uint64_t *)Address, utils::bitCast<uint64_t>(Val), Ordering, MemScope));
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<!utils::is_floating_point_v<V>, V>
min(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_min(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<utils::is_same_v<V, float>, V>
min(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
if (Val >= 0)
return utils::bitCast<float>(min(
(int32_t *)Address, utils::bitCast<int32_t>(Val), Ordering, MemScope));
return utils::bitCast<float>(max(
(uint32_t *)Address, utils::bitCast<uint32_t>(Val), Ordering, MemScope));
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
utils::enable_if_t<utils::is_same_v<V, double>, V>
min(Ty *Address, utils::remove_addrspace_t<Ty> Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
if (Val >= 0)
return utils::bitCast<double>(min(
(int64_t *)Address, utils::bitCast<int64_t>(Val), Ordering, MemScope));
return utils::bitCast<double>(max(
(uint64_t *)Address, utils::bitCast<uint64_t>(Val), Ordering, MemScope));
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V bit_or(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_or(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V bit_and(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_and(Address, Val, Ordering, MemScope);
}
template <typename Ty, typename V = utils::remove_addrspace_t<Ty>>
V bit_xor(Ty *Address, V Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
return __scoped_atomic_fetch_xor(Address, Val, Ordering, MemScope);
}
static inline uint32_t
atomicExchange(uint32_t *Address, uint32_t Val, atomic::OrderingTy Ordering,
MemScopeTy MemScope = MemScopeTy::device) {
uint32_t R;
__scoped_atomic_exchange(Address, &Val, &R, Ordering, MemScope);
return R;
}
}
namespace synchronize {
void init(bool IsSPMD);
void warp(LaneMaskTy Mask);
void threads(atomic::OrderingTy Ordering);
[[gnu::noinline, omp::assume("ompx_aligned_barrier")]] void
threadsAligned(atomic::OrderingTy Ordering);
}
namespace fence {
void team(atomic::OrderingTy Ordering);
void kernel(atomic::OrderingTy Ordering);
void system(atomic::OrderingTy Ordering);
}
}
#endif