1#ifndef INFINI_OPS_DISPATCHER_H_
2#define INFINI_OPS_DISPATCHER_H_
9#include "common/traits.h"
22template <
typename ValueType,
typename Functor,
typename... Args,
auto head,
25 std::string_view context_str, List<head, tail...>,
27 using ReturnType =
decltype(std::forward<Functor>(func)(
28 ValueTag<static_cast<ValueType>(head)>{}, std::forward<Args>(args)...));
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)...),
38 (value ==
static_cast<ValueType
>(head)
39 ? (std::forward<Functor>(func)(
40 ValueTag<head>{}, std::forward<Args>(args)...),
46 std::cerr <<
"dispatch error (void): value " <<
static_cast<int>(value)
47 <<
" not supported in the context: " << context_str <<
"\n";
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)...)),
60 (value ==
static_cast<ValueType
>(head)
61 ? (result.emplace(std::forward<Functor>(func)(
62 ValueTag<head>{}, std::forward<Args>(args)...)),
70 std::cerr <<
"dispatch error (non-void): value " <<
static_cast<int>(value)
71 <<
" not supported in the context: " << context_str <<
"\n";
79template <
typename ValueType,
typename Functor,
typename FilteredList,
81struct DispatchFuncUnwrap;
83template <
typename ValueType,
typename Functor,
auto head,
auto... tail,
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) {
90 List<head, tail...>{}, std::forward<Args>(args)...);
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,
100 std::cerr <<
"dispatch error: no allowed values registered for value "
101 <<
static_cast<int64_t
>(value)
102 <<
" in the context: " << context_str <<
"\n";
110template <
typename ValueType, ValueType... all_values,
typename Functor,
113 std::string_view context_str =
"", Args&&... args) {
114 using FilteredPack =
typename Filter<Functor, std::tuple<Args...>, List<>,
115 all_values...>::type;
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)...);
126template <
typename Functor,
typename... Args,
auto... items>
128 Functor&& func, std::string_view ,
129 List<items...>, Args&&... args) {
130 return std::forward<Functor>(func)(List<items...>{},
131 std::forward<Args>(args)...);
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...>,
143template <
typename RestListsPack,
typename Functor,
auto... items>
146template <
typename... RestLists,
typename Functor,
auto... items>
153 template <
auto val,
typename... Args>
155 return DispatchFunc<RestLists...>(values, next_index, func, context_str,
156 List<items..., val>{},
157 std::forward<Args>(args)...);
161template <
typename RestListsPack,
typename Functor,
typename... Args,
162 auto... items,
auto... allowed>
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)...>;
171 values, index + 1, func, context_str};
174 static_cast<EnumType
>(values.at(index)), adapter, context_str,
175 std::forward<Args>(args)...);
179template <
typename FirstList,
typename... RestLists,
typename Functor,
180 typename... Args,
auto... items>
182 Functor&& func, std::string_view context_str, List<items...>,
185 values, index, func, context_str, List<items...>{}, FirstList{},
186 std::forward<Args>(args)...);
198template <Device::Type kDev,
typename Functor>
199struct DataTypeAdapter {
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)...);
209template <Device::Type kDev,
typename Functor>
210struct DataTypeMultiAdapter {
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)...);
220template <
typename Functor>
221struct DeviceAdapter {
224 template <
auto dev,
typename... Args>
225 auto operator()(ValueTag<dev>, Args&&... args)
const {
226 return func(ValueTag<dev>{}, std::forward<Args>(args)...);
230template <
typename Functor>
231struct DeviceMultiAdapter {
234 template <
auto... devs,
typename... Args>
235 auto operator()(List<devs...>, Args&&... args)
const {
236 return func(ValueTag<devs>{}..., std::forward<Args>(args)...);
243template <Device::Type kDev, DataType... allowed_dtypes,
typename Functor,
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)...);
253template <Device::Type kDev,
typename... Lists,
typename Functor,
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));
260 detail::DataTypeMultiAdapter<kDev, std::remove_reference_t<Functor>> adapter{
262 return DispatchFunc<Lists...>(v, 0, adapter, context_str, List<>{},
263 std::forward<Args>(args)...);
267template <
auto... allowed_devices,
typename Functor,
typename... Args>
269 std::string_view context_str =
"", Args&&... args) {
270 detail::DeviceAdapter<std::remove_reference_t<Functor>> adapter{func};
272 static_cast<Device::Type
>(allowed_devices)...>(
273 device, adapter, context_str, std::forward<Args>(args)...);
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));
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)...);
288template <
typename ValueType,
typename Functor,
typename... Args,
auto... items>
290 std::string_view context_str, List<items...>,
292 return DispatchFunc<static_cast<std::decay_t<ValueType>>(items)...>(
293 value, std::forward<Functor>(func), context_str,
294 std::forward<Args>(args)...);
297template <Device::Type kDev,
typename ValueType,
typename Functor,
298 typename... Args,
auto... items>
300 std::string_view context_str, List<items...>,
302 return DispatchFunc<kDev, static_cast<std::decay_t<ValueType>>(items)...>(
303 value, std::forward<Functor>(func), context_str,
304 std::forward<Args>(args)...);
308template <
typename ListType,
typename ValueType,
typename Functor,
310 typename = std::enable_if_t<IsListType<ListType>::value>>
312 std::string_view context_str =
"", Args&&... args) {
314 context_str, ListType{},
315 std::forward<Args>(args)...);
319template <Device::Type kDev,
typename ListType,
typename ValueType,
320 typename Functor,
typename... Args,
321 typename = std::enable_if_t<IsListType<ListType>::value>>
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)...);
330template <
typename... Lists,
typename Functor,
typename... Args>
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)...);
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
size_t next_index
Definition dispatcher.h:149
Functor & func
Definition dispatcher.h:150
std::string_view context_str
Definition dispatcher.h:151
Definition dispatcher.h:144