13#ifndef PROTEUS_JIT_INTERFACE_H
14#define PROTEUS_JIT_INTERFACE_H
26__proteus_register_lambda_runtime_constant(int32_t
Type, int32_t
Pos,
32__proteus_finalize_register(uint64_t Tag)
noexcept;
36#if defined(__CUDACC__) || defined(__HIP__)
37#define PROTEUS_HOST_DEVICE __host__ __device__
39#define PROTEUS_HOST_DEVICE
42template <
typename T>
__attribute__((noinline))
void jit_arg(T V)
noexcept;
43#if defined(__CUDACC__) || defined(__HIP__)
45__attribute__((noinline)) __device__
void jit_arg(T V)
noexcept;
50jit_array(T V, [[maybe_unused]]
size_t NumElts,
52 typename std::remove_pointer<T>::type
Velem = 0) noexcept;
53#if defined(__CUDACC__) || defined(__HIP__)
56jit_array(T V, [[maybe_unused]]
size_t NumElts,
58 typename std::remove_pointer<T>::type
Velem = 0) noexcept;
63std::enable_if_t<std::is_trivially_copyable_v<std::remove_pointer_t<T>>,
void>
64jit_object(T *V,
size_t Size =
sizeof(std::remove_pointer_t<T>))
noexcept;
66#if defined(__CUDACC__) || defined(__HIP__)
69 std::is_trivially_copyable_v<std::remove_pointer_t<T>>,
void>
70jit_object(T *V,
size_t Size =
sizeof(T))
noexcept;
75std::enable_if_t<!std::is_pointer_v<T> &&
76 std::is_trivially_copyable_v<std::remove_reference_t<T>>,
78jit_object(T &V,
size_t Size =
sizeof(std::remove_reference_t<T>))
noexcept;
80#if defined(__CUDACC__) || defined(__HIP__)
83 !std::is_pointer_v<T> &&
84 std::is_trivially_copyable_v<std::remove_reference_t<T>>,
86jit_object(T &V,
size_t Size =
sizeof(T))
noexcept;
91constexpr std::uint64_t
fnv1a64(
const char *s) {
92 std::uint64_t h = 14695981039346656037ull;
94 h ^= (
unsigned char)(*s);
95 h *= 1099511628211ull;
101template <
class Lambda>
constexpr std::uint64_t
functor_id() {
102 return fnv1a64(__PRETTY_FUNCTION__);
110 template <
typename...
Args>
113 operator()(
Args &&...args) const noexcept {
114 return lambda(std::forward<Args>(args)...);
120template <u
int64_t FunctorID,
typename Lambda>
124template <std::u
int64_t FunctorId,
typename L>
127 std::forward<L>(lambda)};
130template <u
int64_t ID,
typename T>
132__attribute__((annotate(
"proteus.register_call_impl", ID)))
auto
138 using LambdaType = std::decay_t<T>;
139 LambdaType local = t;
141 auto result = tag_functor<ID>(std::forward<T>(t));
149register_lambda(L &&lambda) noexcept {
151 ::proteus::detail::functor_id<std::decay_t<L>>()>(
152 std::forward<L>(lambda));
154 return registered_lambda;
158static __attribute__((noinline)) T jit_variable(T V)
noexcept {
162#if defined(__CUDACC__) || defined(__HIP__)
165template <
typename T,
size_t MAXN,
int UniqueID = 0>
167shared_array([[maybe_unused]]
size_t N,
168 [[maybe_unused]]
size_t ElemSize =
sizeof(T)) {
169 alignas(T)
static __shared__
char shmem[
sizeof(T) * MAXN];
170 return reinterpret_cast<T *
>(shmem);
uint32_t int32_t Type
Definition CompilerInterfaceDevice.cpp:98
__attribute__((used)) void __proteus_register_fatbinary(void *Handle
Definition CompilerInterfaceDevice.cpp:49
char int void ** Args
Definition CompilerInterfaceHost.cpp:23
__attribute__((used)) void __proteus_register_lambda_runtime_constant(int32_t Type
#define PROTEUS_HOST_DEVICE
Definition JitInterface.h:39
int32_t int32_t const void * ValuePtr
Definition JitInterface.h:27
int32_t int32_t Offset
Definition JitInterface.h:27
int32_t int32_t const void uint64_t functor_id
Definition JitInterface.h:28
void __proteus_take_address(void const *) noexcept
int32_t Pos
Definition JitInterface.h:26
PROTEUS_HOST_DEVICE auto tag_functor(L &&lambda)
Definition JitInterface.h:125
static ID auto __register_lambda_impl(T &&t) noexcept
Definition JitInterface.h:133
constexpr std::uint64_t fnv1a64(const char *s)
Definition JitInterface.h:91
constexpr std::uint64_t functor_id()
Definition JitInterface.h:101
Definition MemoryCache.h:27
size_t std::remove_pointer< T >::type Velem
Definition JitInterface.h:52
size_t NumElts
Definition JitInterface.h:50
__attribute__((noinline)) void jit_arg(T V) noexcept
Definition JitInterface.h:105
Lambda LambdaType
Definition JitInterface.h:106
PROTEUS_HOST_DEVICE __attribute__((annotate("jit"))) __attribute__((annotate("proteus.wrapper_call"
LambdaType lambda
Definition JitInterface.h:108
static constexpr std::uint64_t functor_id
Definition JitInterface.h:107
Definition JitInterface.h:118