PXL
direction.hpp
Go to the documentation of this file.
1 // SPDX-License-Identifier: Apache-2.0
2 // Copyright 2026 XCENA Inc.
3 
4 #pragma once
5 
6 #include <cstddef>
7 #include <tuple>
8 #include <type_traits>
9 #include <utility>
10 #include <vector>
11 
12 namespace pxl
13 {
14 
15 template <typename T>
16 class NDArray;
17 
18 enum class ArgDir : unsigned char
19 {
20  InOut,
21  Input,
22  Output,
23 };
24 
26 {
27  // Empty means the legacy InOut policy for every argument.
28  std::vector<ArgDir> args;
29 };
30 
31 namespace detail
32 {
33 
34 template <typename T>
36 
37 template <typename R, typename... Args>
38 struct KernelFunctionTraits<R (*)(Args...)>
39 {
40  static constexpr std::size_t arity = sizeof...(Args);
41  using parameters = std::tuple<Args...>;
42 };
43 
44 template <typename R, typename... Args>
45 struct KernelFunctionTraits<R (*)(Args...) noexcept> : KernelFunctionTraits<R (*)(Args...)>
46 {
47 };
48 
49 template <auto Func>
50 inline constexpr std::size_t KernelArityV = KernelFunctionTraits<decltype(Func)>::arity;
51 
52 template <auto Func, std::size_t Index>
53 using KernelParamT =
54  std::tuple_element_t<Index, typename KernelFunctionTraits<decltype(Func)>::parameters>;
55 
56 template <typename T, typename = void>
58 {
59  static constexpr ArgDir value = ArgDir::InOut;
60 };
61 
62 template <typename T>
64  T, std::void_t<typename T::element_type, typename T::kernel_buffer_direction>>
65 {
66 private:
67  using direction_tag = typename T::kernel_buffer_direction;
68  static constexpr bool reads = direction_tag::reads_existing_data;
69  static constexpr bool writes = direction_tag::writes_result;
70 
71 public:
72  static_assert(reads || writes, "A directional kernel buffer needs read or write access");
73  static constexpr ArgDir value = reads ? (writes ? ArgDir::InOut : ArgDir::Input)
75 };
76 
77 template <typename T>
78 inline constexpr ArgDir KernelParamDirectionV =
80 
81 template <typename T, typename = void>
82 struct IsDirectionalKernelParam : std::false_type
83 {
84 };
85 
86 template <typename T>
88  T, std::void_t<typename T::element_type, typename T::kernel_buffer_direction>>
89  : std::true_type
90 {
91 };
92 
93 template <typename T>
94 inline constexpr bool IsDirectionalKernelParamV =
96 
97 template <typename T>
99 {
100  static constexpr bool valid = false;
101  using element_type = void;
102  using pointer = void*;
103 };
104 
105 template <typename T>
107 {
108  static constexpr bool valid = true;
109  using element_type = T;
110  using pointer = T*;
111 };
112 
113 template <typename T>
115 {
116  static constexpr bool valid = true;
117  using element_type = T;
118  using pointer = T*;
119 };
120 
121 template <typename Param, typename Actual, bool = IsDirectionalKernelParamV<Param>>
122 struct KernelArgumentCompatible : std::true_type
123 {
124 };
125 
126 template <typename Param, typename Actual>
127 struct KernelArgumentCompatible<Param, Actual, true>
128 {
129 private:
130  using parameter_type = std::remove_cv_t<std::remove_reference_t<Param>>;
131  using parameter_element = typename parameter_type::element_type;
132  using direction_tag = typename parameter_type::kernel_buffer_direction;
134  using actual_element = typename actual_buffer::element_type;
135  static constexpr bool writes = direction_tag::writes_result;
136  using wire_element = std::conditional_t<writes, parameter_element,
137  std::add_const_t<parameter_element>>;
138 
139 public:
140  static constexpr bool value =
141  !std::is_reference_v<Param> &&
142  actual_buffer::valid &&
143  !std::is_void_v<std::remove_cv_t<parameter_element>> &&
144  !std::is_function_v<parameter_element> &&
145  (!writes || !std::is_const_v<parameter_element>) &&
146  std::is_same_v<std::remove_cv_t<actual_element>,
147  std::remove_cv_t<parameter_element>> &&
148  std::is_convertible_v<typename actual_buffer::pointer, wire_element*>;
149 };
150 
151 template <typename Param, typename Actual>
152 inline constexpr bool KernelArgumentCompatibleV =
154 
155 template <auto Func, typename ActualTuple, std::size_t... Indices>
156 constexpr bool KernelArgumentsCompatible(std::index_sequence<Indices...>)
157 {
160  std::tuple_element_t<Indices, ActualTuple>> &&
161  ...);
162 }
163 
164 template <auto Func, typename... Actual>
166 {
167  if constexpr (sizeof...(Actual) != KernelArityV<Func>)
168  {
169  return false;
170  }
171  else
172  {
173  return KernelArgumentsCompatible<Func, std::tuple<Actual...>>(
174  std::make_index_sequence<KernelArityV<Func>>{});
175  }
176 }
177 
178 template <auto Func, typename... Actual>
179 inline constexpr bool KernelArgumentsCompatibleV =
180  KernelArgumentsCompatible<Func, Actual...>();
181 
182 template <auto Func, typename... Actual>
184  std::enable_if_t<KernelArgumentsCompatibleV<Func, Actual...>, int>;
185 
186 template <auto Func, std::size_t... Indices>
187 ArgDirections MakeKernelDirections(std::index_sequence<Indices...>)
188 {
189  ArgDirections directions;
191  {
192  return directions;
193  }
194  directions.args = {KernelParamDirectionV<KernelParamT<Func, Indices>>...};
195  return directions;
196 }
197 
198 template <auto Func>
200 {
201  return MakeKernelDirections<Func>(std::make_index_sequence<KernelArityV<Func>>{});
202 }
203 
204 } // namespace detail
205 
206 namespace impl
207 {
208 
209 inline ArgDir directionAt(const std::vector<ArgDir>& directions, std::size_t index) noexcept
210 {
211  return index < directions.size() ? directions[index] : ArgDir::InOut;
212 }
213 
214 } // namespace impl
215 
216 } // namespace pxl
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
Definition: ndarray.hpp:25
ArgDirections MakeKernelDirections(std::index_sequence< Indices... >)
Definition: direction.hpp:187
constexpr bool KernelArgumentCompatibleV
Definition: direction.hpp:152
constexpr ArgDir KernelParamDirectionV
Definition: direction.hpp:78
constexpr std::size_t KernelArityV
Definition: direction.hpp:50
constexpr bool IsDirectionalKernelParamV
Definition: direction.hpp:94
constexpr bool KernelArgumentsCompatibleV
Definition: direction.hpp:179
std::enable_if_t< KernelArgumentsCompatibleV< Func, Actual... >, int > EnableCompatibleKernelArguments
Definition: direction.hpp:184
constexpr bool KernelArgumentsCompatible(std::index_sequence< Indices... >)
Definition: direction.hpp:156
std::tuple_element_t< Index, typename KernelFunctionTraits< decltype(Func)>::parameters > KernelParamT
Definition: direction.hpp:54
Definition: config.hpp:11
ArgDir
Definition: direction.hpp:19
std::vector< ArgDir > args
Definition: direction.hpp:28
static constexpr bool valid
Definition: direction.hpp:100
static constexpr ArgDir value
Definition: direction.hpp:59