InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
dispatcher.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_DISPATCHER_H_
2#define INFINI_OPS_DISPATCHER_H_
3
4#include <iostream>
5#include <optional>
6#include <string_view>
7#include <vector>
8
9#include "common/traits.h"
10#include "data_type.h"
11#include "device.h"
12
13namespace infini::ops {
14
15// -----------------------------------------------------------------------------
16// Core Generic Runtime Dispatchers
17// -----------------------------------------------------------------------------
18
19namespace detail {
20
21// Implements the dispatch body over a resolved `List<head, tail...>`.
22template <typename ValueType, typename Functor, typename... Args, auto head,
23 auto... tail>
24auto DispatchFuncImpl(ValueType value, Functor&& func,
25 std::string_view context_str, List<head, tail...>,
26 Args&&... args) {
27 using ReturnType = decltype(std::forward<Functor>(func)(
28 ValueTag<static_cast<ValueType>(head)>{}, std::forward<Args>(args)...));
29
30 // Path for void functions.
31 if constexpr (std::is_void_v<ReturnType>) {
32 bool handled = ((value == static_cast<ValueType>(tail)
33 ? (std::forward<Functor>(func)(
34 ValueTag<tail>{}, std::forward<Args>(args)...),
35 true)
36 : false) ||
37 ... ||
38 (value == static_cast<ValueType>(head)
39 ? (std::forward<Functor>(func)(
40 ValueTag<head>{}, std::forward<Args>(args)...),
41 true)
42 : false));
43
44 if (!handled) {
45 // TODO(lzm): change to logging.
46 std::cerr << "dispatch error (void): value " << static_cast<int>(value)
47 << " not supported in the context: " << context_str << "\n";
48 std::abort();
49 }
50 }
51 // Path for non-void functions.
52 else {
53 std::optional<ReturnType> result;
54 bool handled = ((value == static_cast<ValueType>(tail)
55 ? (result.emplace(std::forward<Functor>(func)(
56 ValueTag<tail>{}, std::forward<Args>(args)...)),
57 true)
58 : false) ||
59 ... ||
60 (value == static_cast<ValueType>(head)
61 ? (result.emplace(std::forward<Functor>(func)(
62 ValueTag<head>{}, std::forward<Args>(args)...)),
63 true)
64 : false));
65
66 if (handled) {
67 return *result;
68 }
69 // TODO(lzm): change to logging.
70 std::cerr << "dispatch error (non-void): value " << static_cast<int>(value)
71 << " not supported in the context: " << context_str << "\n";
72 std::abort();
73 return ReturnType{};
74 }
75}
76
77// Deduces `head`/`tail` from a `List` type via partial specialization,
78// then forwards to `DispatchFuncImpl`.
79template <typename ValueType, typename Functor, typename FilteredList,
80 typename ArgsTuple>
81struct DispatchFuncUnwrap;
82
83template <typename ValueType, typename Functor, auto head, auto... tail,
84 typename... Args>
85struct DispatchFuncUnwrap<ValueType, Functor, List<head, tail...>,
86 std::tuple<Args...>> {
87 static auto Call(ValueType value, Functor&& func,
88 std::string_view context_str, Args&&... args) {
89 return DispatchFuncImpl(value, std::forward<Functor>(func), context_str,
90 List<head, tail...>{}, std::forward<Args>(args)...);
91 }
92};
93
94// Empty-list specialization
95template <typename ValueType, typename Functor, typename... Args>
96struct DispatchFuncUnwrap<ValueType, Functor, List<>, std::tuple<Args...>> {
97 static auto Call(ValueType value, Functor&&, std::string_view context_str,
98 Args&&...) {
99 // TODO(lzm): change to logging.
100 std::cerr << "dispatch error: no allowed values registered for value "
101 << static_cast<int64_t>(value)
102 << " in the context: " << context_str << "\n";
103 std::abort();
104 }
105};
106
107} // namespace detail
108
109// (Single Dispatch) Dispatches a runtime value to a compile-time functor.
110template <typename ValueType, ValueType... all_values, typename Functor,
111 typename... Args>
112auto DispatchFunc(ValueType value, Functor&& func,
113 std::string_view context_str = "", Args&&... args) {
114 using FilteredPack = typename Filter<Functor, std::tuple<Args...>, List<>,
115 all_values...>::type;
116
117 return detail::DispatchFuncUnwrap<
118 ValueType, Functor, FilteredPack,
119 std::tuple<Args...>>::Call(value, std::forward<Functor>(func),
120 context_str, std::forward<Args>(args)...);
121}
122
123// (Multi-Dispatch) Dispatches a vector of runtime values to a compile-time
124// functor.
125// Base Case: All Dimensions Resolved
126template <typename Functor, typename... Args, auto... items>
127auto DispatchFunc(const std::vector<int64_t>& values, size_t /*index*/,
128 Functor&& func, std::string_view /*context_str*/,
129 List<items...>, Args&&... args) {
130 return std::forward<Functor>(func)(List<items...>{},
131 std::forward<Args>(args)...);
132}
133
134// Forward declaration of the recursive multi-dispatch overload.
135template <typename FirstList, typename... RestLists, typename Functor,
136 typename... Args, auto... items>
137auto DispatchFunc(const std::vector<int64_t>& values, size_t index,
138 Functor&& func, std::string_view context_str, List<items...>,
139 Args&&... args);
140
141// Adapter used in the recursive multi-dispatch case: given a resolved value
142// `val` recurse into the next dimension.
143template <typename RestListsPack, typename Functor, auto... items>
145
146template <typename... RestLists, typename Functor, auto... items>
147struct MultiDispatchRecurseAdapter<TypePack<RestLists...>, Functor, items...> {
148 const std::vector<int64_t>& values;
150 Functor& func;
151 std::string_view context_str;
152
153 template <auto val, typename... Args>
154 auto operator()(ValueTag<val>, Args&&... args) const {
155 return DispatchFunc<RestLists...>(values, next_index, func, context_str,
156 List<items..., val>{},
157 std::forward<Args>(args)...);
158 }
159};
160
161template <typename RestListsPack, typename Functor, typename... Args,
162 auto... items, auto... allowed>
163auto MultiDispatchFirstDim(const std::vector<int64_t>& values, size_t index,
164 Functor& func, std::string_view context_str,
165 List<items...>, List<allowed...>, Args&&... args) {
166 static_assert(sizeof...(allowed) > 0,
167 "`DispatchFunc` dimension list is empty");
168 using EnumType = std::common_type_t<decltype(allowed)...>;
169
170 MultiDispatchRecurseAdapter<RestListsPack, Functor, items...> adapter{
171 values, index + 1, func, context_str};
172
173 return DispatchFunc<EnumType, allowed...>(
174 static_cast<EnumType>(values.at(index)), adapter, context_str,
175 std::forward<Args>(args)...);
176}
177
178// (Multi-Dispatch) Recursive Case
179template <typename FirstList, typename... RestLists, typename Functor,
180 typename... Args, auto... items>
181auto DispatchFunc(const std::vector<int64_t>& values, size_t index,
182 Functor&& func, std::string_view context_str, List<items...>,
183 Args&&... args) {
184 return MultiDispatchFirstDim<TypePack<RestLists...>>(
185 values, index, func, context_str, List<items...>{}, FirstList{},
186 std::forward<Args>(args)...);
187}
188
189// -----------------------------------------------------------------------------
190// High-Level Specialized Dispatchers
191// -----------------------------------------------------------------------------
192// These provide cleaner and more convenient APIs for common InfiniOps types.
193
194namespace detail {
195
196// Bridges the generic value dispatch layer to the `DataType`-specific type
197// dispatch layer.
198template <Device::Type kDev, typename Functor>
199struct DataTypeAdapter {
200 Functor& func;
201
202 template <auto dtype, typename... Args>
203 auto operator()(ValueTag<dtype>, Args&&... args) const {
204 using T = TypeMapType<kDev, static_cast<DataType>(dtype)>;
205 return func(TypeTag<T>{}, std::forward<Args>(args)...);
206 }
207};
208
209template <Device::Type kDev, typename Functor>
210struct DataTypeMultiAdapter {
211 Functor& func;
212
213 template <auto... dtypes, typename... Args>
214 auto operator()(List<dtypes...>, Args&&... args) const {
215 return func(TypeTag<TypeMapType<kDev, static_cast<DataType>(dtypes)>>{}...,
216 std::forward<Args>(args)...);
217 }
218};
219
220template <typename Functor>
221struct DeviceAdapter {
222 Functor& func;
223
224 template <auto dev, typename... Args>
225 auto operator()(ValueTag<dev>, Args&&... args) const {
226 return func(ValueTag<dev>{}, std::forward<Args>(args)...);
227 }
228};
229
230template <typename Functor>
231struct DeviceMultiAdapter {
232 Functor& func;
233
234 template <auto... devs, typename... Args>
235 auto operator()(List<devs...>, Args&&... args) const {
236 return func(ValueTag<devs>{}..., std::forward<Args>(args)...);
237 }
238};
239
240} // namespace detail
241
242// `DataType` Dispatch
243template <Device::Type kDev, DataType... allowed_dtypes, typename Functor,
244 typename... Args>
245auto DispatchFunc(DataType dtype, Functor&& func,
246 std::string_view context_str = "", Args&&... args) {
247 detail::DataTypeAdapter<kDev, std::remove_reference_t<Functor>> adapter{func};
248 return DispatchFunc<DataType, allowed_dtypes...>(dtype, adapter, context_str,
249 std::forward<Args>(args)...);
250}
251
252// `DataType` Multi-Dispatch
253template <Device::Type kDev, typename... Lists, typename Functor,
254 typename... Args>
255auto DispatchFunc(std::initializer_list<DataType> dtypes, Functor&& func,
256 std::string_view context_str = "", Args&&... args) {
257 std::vector<int64_t> v;
258 for (auto d : dtypes) v.push_back(static_cast<int64_t>(d));
259
260 detail::DataTypeMultiAdapter<kDev, std::remove_reference_t<Functor>> adapter{
261 func};
262 return DispatchFunc<Lists...>(v, 0, adapter, context_str, List<>{},
263 std::forward<Args>(args)...);
264}
265
266// `Device` Dispatch
267template <auto... allowed_devices, typename Functor, typename... Args>
268auto DispatchFunc(Device::Type device, Functor&& func,
269 std::string_view context_str = "", Args&&... args) {
270 detail::DeviceAdapter<std::remove_reference_t<Functor>> adapter{func};
271 return DispatchFunc<Device::Type,
272 static_cast<Device::Type>(allowed_devices)...>(
273 device, adapter, context_str, std::forward<Args>(args)...);
274}
275
276// `Device` Multi-Dispatch
277template <typename... Lists, typename Functor, typename... Args>
278auto DispatchFunc(std::initializer_list<Device::Type> devices, Functor&& func,
279 std::string_view context_str = "", Args&&... args) {
280 std::vector<int64_t> v;
281 for (auto d : devices) v.push_back(static_cast<int64_t>(d));
282
283 detail::DeviceMultiAdapter<std::remove_reference_t<Functor>> adapter{func};
284 return DispatchFunc<Lists...>(v, 0, adapter, context_str, List<>{},
285 std::forward<Args>(args)...);
286}
287
288template <typename ValueType, typename Functor, typename... Args, auto... items>
289auto DispatchFuncListAliasImpl(ValueType value, Functor&& func,
290 std::string_view context_str, List<items...>,
291 Args&&... args) {
292 return DispatchFunc<static_cast<std::decay_t<ValueType>>(items)...>(
293 value, std::forward<Functor>(func), context_str,
294 std::forward<Args>(args)...);
295}
296
297template <Device::Type kDev, typename ValueType, typename Functor,
298 typename... Args, auto... items>
299auto DispatchFuncListAliasImpl(ValueType value, Functor&& func,
300 std::string_view context_str, List<items...>,
301 Args&&... args) {
302 return DispatchFunc<kDev, static_cast<std::decay_t<ValueType>>(items)...>(
303 value, std::forward<Functor>(func), context_str,
304 std::forward<Args>(args)...);
305}
306
307// Interface for Generic `List` Aliases (for non-DataType dispatch, e.g. Device)
308template <typename ListType, typename ValueType, typename Functor,
309 typename... Args,
310 typename = std::enable_if_t<IsListType<ListType>::value>>
311auto DispatchFunc(ValueType value, Functor&& func,
312 std::string_view context_str = "", Args&&... args) {
313 return DispatchFuncListAliasImpl(value, std::forward<Functor>(func),
314 context_str, ListType{},
315 std::forward<Args>(args)...);
316}
317
318// Interface for Generic `List` Aliases (for DataType dispatch with device type)
319template <Device::Type kDev, typename ListType, typename ValueType,
320 typename Functor, typename... Args,
321 typename = std::enable_if_t<IsListType<ListType>::value>>
322auto DispatchFunc(ValueType value, Functor&& func,
323 std::string_view context_str = "", Args&&... args) {
324 return DispatchFuncListAliasImpl<kDev>(value, std::forward<Functor>(func),
325 context_str, ListType{},
326 std::forward<Args>(args)...);
327}
328
329// Interface for Any `int64_t`-Convertible Types
330template <typename... Lists, typename Functor, typename... Args>
331auto DispatchFunc(std::initializer_list<int64_t> keys, Functor&& func,
332 std::string_view context_str = "", Args&&... args) {
333 std::vector<int64_t> v_keys(keys);
334 return DispatchFunc<Lists...>(v_keys, 0, std::forward<Functor>(func),
335 context_str, List<>{},
336 std::forward<Args>(args)...);
337}
338
339} // namespace infini::ops
340
341#endif
auto DispatchFuncImpl(ValueType value, Functor &&func, std::string_view context_str, List< head, tail... >, Args &&... args)
Definition dispatcher.h:24
Definition generated/include/operator.h:28
auto MultiDispatchFirstDim(const std::vector< int64_t > &values, size_t index, Functor &func, std::string_view context_str, List< items... >, List< allowed... >, Args &&... args)
Definition dispatcher.h:163
auto DispatchFuncListAliasImpl(ValueType value, Functor &&func, std::string_view context_str, List< items... >, Args &&... args)
Definition dispatcher.h:289
infini::rt::TypeMapType< dev, dtype > TypeMapType
Definition data_type.h:24
auto DispatchFunc(ValueType value, Functor &&func, std::string_view context_str="", Args &&... args)
Definition dispatcher.h:112
const std::vector< int64_t > & values
Definition dispatcher.h:148
auto operator()(ValueTag< val >, Args &&... args) const
Definition dispatcher.h:154
Definition dispatcher.h:144