// ROCm stand-in for libcu++'s <cuda/atomic> (ROCm ships no libcu++). On the HIP include path only.
//
// Covers the slice the kernels use: cuda::atomic_ref<int, cuda::thread_scope_device> with load, store and
// fetch_add, for group_barrier in ptx.cuh and the MoE scheduler in quant/exl3_moe_kernel.cuh. Anything outside
// that slice fails to compile rather than silently taking a weaker ordering.

#pragma once
#include <hip/hip_runtime.h>

namespace cuda {

enum thread_scope
{
    thread_scope_system = __HIP_MEMORY_SCOPE_SYSTEM,
    thread_scope_device = __HIP_MEMORY_SCOPE_AGENT,
    thread_scope_block  = __HIP_MEMORY_SCOPE_WORKGROUP,
    thread_scope_thread = __HIP_MEMORY_SCOPE_SINGLETHREAD,
};

inline constexpr int memory_order_relaxed = __ATOMIC_RELAXED;
inline constexpr int memory_order_acquire = __ATOMIC_ACQUIRE;
inline constexpr int memory_order_release = __ATOMIC_RELEASE;
inline constexpr int memory_order_acq_rel = __ATOMIC_ACQ_REL;
inline constexpr int memory_order_seq_cst = __ATOMIC_SEQ_CST;

template <typename T, thread_scope Scope = thread_scope_device>
class atomic_ref
{
public:
    __device__ explicit atomic_ref(T& ref) : ptr(&ref) {}
    __device__ T load(int order = memory_order_seq_cst) const { return __hip_atomic_load(ptr, order, Scope); }
    __device__ void store(T v, int order = memory_order_seq_cst) const { __hip_atomic_store(ptr, v, order, Scope); }
    __device__ T fetch_add(T v, int order = memory_order_seq_cst) const { return __hip_atomic_fetch_add(ptr, v, order, Scope); }

private:
    T* ptr;
};

}  // namespace cuda
