Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
JitInterface.h
Go to the documentation of this file.
1//===-- jit.h -- user interface to Proteus JIT library --===//
2//
3// Part of the Proteus Project, under the Apache License v2.0 with LLVM
4// Exceptions. See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9//===----------------------------------------------------------------------===//
10
11// NOLINTBEGIN(readability-identifier-naming)
12
13#ifndef PROTEUS_JIT_INTERFACE_H
14#define PROTEUS_JIT_INTERFACE_H
15
17#include "proteus/Init.h"
18
19#include <cassert>
20#include <cstdint>
21#include <cstring>
22#include <type_traits>
23#include <utility>
24
25extern "C" __attribute__((used)) void
26__proteus_register_lambda_runtime_constant(int32_t Type, int32_t Pos,
27 int32_t Offset, const void *ValuePtr,
28 uint64_t functor_id);
29
30extern "C" void __proteus_take_address(void const *) noexcept;
31extern "C" __attribute__((used)) void
32__proteus_finalize_register(uint64_t Tag) noexcept;
33
34namespace proteus {
35
36#if defined(__CUDACC__) || defined(__HIP__)
37#define PROTEUS_HOST_DEVICE __host__ __device__
38#else
39#define PROTEUS_HOST_DEVICE
40#endif
41
42template <typename T> __attribute__((noinline)) void jit_arg(T V) noexcept;
43#if defined(__CUDACC__) || defined(__HIP__)
44template <typename T>
45__attribute__((noinline)) __device__ void jit_arg(T V) noexcept;
46#endif
47
48template <typename T>
49__attribute__((noinline)) void
50jit_array(T V, [[maybe_unused]] size_t NumElts,
51 [[maybe_unused]]
52 typename std::remove_pointer<T>::type Velem = 0) noexcept;
53#if defined(__CUDACC__) || defined(__HIP__)
54template <typename T>
55__attribute__((noinline)) __device__ void
56jit_array(T V, [[maybe_unused]] size_t NumElts,
57 [[maybe_unused]]
58 typename std::remove_pointer<T>::type Velem = 0) noexcept;
59#endif
60
61template <typename T>
62__attribute__((noinline))
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;
65
66#if defined(__CUDACC__) || defined(__HIP__)
67template <typename T>
68__attribute__((noinline)) __device__ std::enable_if_t<
69 std::is_trivially_copyable_v<std::remove_pointer_t<T>>, void>
70jit_object(T *V, size_t Size = sizeof(T)) noexcept;
71#endif
72
73template <typename T>
74__attribute__((noinline))
75std::enable_if_t<!std::is_pointer_v<T> &&
76 std::is_trivially_copyable_v<std::remove_reference_t<T>>,
77 void>
78jit_object(T &V, size_t Size = sizeof(std::remove_reference_t<T>)) noexcept;
79
80#if defined(__CUDACC__) || defined(__HIP__)
81template <typename T>
82__attribute__((noinline)) __device__ std::enable_if_t<
83 !std::is_pointer_v<T> &&
84 std::is_trivially_copyable_v<std::remove_reference_t<T>>,
85 void>
86jit_object(T &V, size_t Size = sizeof(T)) noexcept;
87#endif
88
89namespace detail {
90// todo: use LLVM hashing?
91constexpr std::uint64_t fnv1a64(const char *s) {
92 std::uint64_t h = 14695981039346656037ull;
93 for (; *s; ++s) {
94 h ^= (unsigned char)(*s);
95 h *= 1099511628211ull;
96 }
97 return h;
98}
99
100// todo: make a test with lambda factory in a header, two separate cpp files
101template <class Lambda> constexpr std::uint64_t functor_id() {
102 return fnv1a64(__PRETTY_FUNCTION__); // includes Lambda + Ctr in the text
103}
104
105template <uint64_t FunctorID, typename Lambda> struct LambdaFunctorWrapper {
106 using LambdaType = Lambda;
107 static constexpr std::uint64_t functor_id = FunctorID;
109
110 template <typename... Args>
112 __attribute__((annotate("proteus.wrapper_call", functor_id))) decltype(auto)
113 operator()(Args &&...args) const noexcept {
114 return lambda(std::forward<Args>(args)...);
115 }
116};
117
118template <typename... T> struct is_lambda_functor_wrapper : std::false_type {};
119
120template <uint64_t FunctorID, typename Lambda>
122 : std::true_type {};
123
124template <std::uint64_t FunctorId, typename L>
125PROTEUS_HOST_DEVICE inline auto tag_functor(L &&lambda) {
127 std::forward<L>(lambda)};
128}
129
130template <uint64_t ID, typename T>
131[[nodiscard]] static __attribute__((noinline))
132__attribute__((annotate("proteus.register_call_impl", ID))) auto
133__register_lambda_impl(T &&t) noexcept {
135 // Force LLVM to generate an AllocaInst of the underlying Clang--generated
136 // anonymous class for T. We remove this after recording the demangled
137 // lambda name.
138 using LambdaType = std::decay_t<T>;
139 LambdaType local = t;
141 auto result = tag_functor<ID>(std::forward<T>(t));
142 return result;
143}
144
145} // namespace detail
146
147template <class L>
148[[nodiscard]] inline auto __attribute__((annotate("proteus.register_call")))
149register_lambda(L &&lambda) noexcept {
150 auto registered_lambda = ::proteus::detail::__register_lambda_impl<
151 ::proteus::detail::functor_id<std::decay_t<L>>()>(
152 std::forward<L>(lambda));
153 __proteus_take_address(&registered_lambda);
154 return registered_lambda;
155}
156
157template <typename T>
158static __attribute__((noinline)) T jit_variable(T V) noexcept {
159 return V;
160}
161
162#if defined(__CUDACC__) || defined(__HIP__)
163// The function needs to be static for RDC compilation to resolve the static
164// shared memory fallback.
165template <typename T, size_t MAXN, int UniqueID = 0>
166static __device__ __attribute__((noinline)) T *
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);
171}
172#endif
173
174} // namespace proteus
175
176#endif
177
178// NOLINTEND(readability-identifier-naming)
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