InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
generated/include/operator.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_OPERATOR_H_
2#define INFINI_OPS_OPERATOR_H_
3
4#include <algorithm>
5#include <atomic>
6#include <cassert>
7#include <chrono>
8#include <cstdio>
9#include <cstdlib>
10#include <iostream>
11#include <limits>
12#include <memory>
13#include <optional>
14#include <string_view>
15#include <tuple>
16#include <type_traits>
17#include <unordered_map>
18#include <utility>
19#include <vector>
20
21#include "config.h"
22#include "dispatcher.h"
23#include "handle.h"
24#include "runtime.h"
25#include "tensor.h"
26#include "tuning.h"
27
29
30struct CacheKey {
31 std::size_t hash;
32
33 std::vector<Tensor> tensors;
34
35 std::size_t scalar_hash;
36
37 template <typename... Args>
38 static CacheKey Build(const Args&... args) {
39 CacheKey key;
40 key.hash = 0;
41 key.scalar_hash = 0;
42 (key.Absorb(args), ...);
43 return key;
44 }
45
46 private:
47 void Absorb(const Tensor& t) {
48 HashCombine(hash, t);
49 tensors.push_back(t);
50 }
51
52 void Absorb(const std::vector<Tensor>& ts) {
53 HashCombine(hash, ts.size());
54 for (const auto& t : ts) {
55 HashCombine(hash, t);
56 tensors.push_back(t);
57 }
58 }
59
60 template <typename T>
61 void Absorb(const T& v) {
62 HashCombine(hash, v);
63 HashCombine(scalar_hash, v);
64 }
65};
66
67template <typename Functor, typename... Args, auto... implementation_indices>
68auto DispatchImplementation(std::size_t implementation_index, Functor&& func,
69 std::string_view context_str,
70 List<implementation_indices...>, Args&&... args) {
71 return DispatchFunc<std::size_t,
72 static_cast<std::size_t>(implementation_indices)...>(
73 implementation_index, std::forward<Functor>(func), context_str,
74 std::forward<Args>(args)...);
75}
76
77template <auto... values>
78std::vector<std::size_t> ListToVector(List<values...>) {
79 return {static_cast<std::size_t>(values)...};
80}
81
82template <typename ValueType, auto... values>
83bool ListContains(ValueType value, List<values...>) {
84 return ((value == static_cast<ValueType>(values)) || ...);
85}
86
87inline void SyncDevice(Device::Type dev_type) {
88 DispatchFunc<ActiveDevices<void>>(
89 dev_type,
90 [](auto device_tag) {
91 constexpr Device::Type kDev = decltype(device_tag)::value;
92 infini::rt::runtime::Runtime<kDev>::DeviceSynchronize();
93 },
94 "SyncDevice");
95}
96
97inline Device::Type FirstDeviceType() { return Device::Type::kCount; }
98
99template <typename First, typename... Rest>
100Device::Type FirstDeviceType(const First& first, const Rest&... rest) {
101 if constexpr (std::is_same_v<std::decay_t<First>, Tensor>) {
102 return first.device().type();
103 } else if constexpr (std::is_same_v<std::decay_t<First>,
104 std::vector<Tensor>>) {
105 return first.empty() ? FirstDeviceType(rest...)
106 : first.front().device().type();
107 } else {
108 return FirstDeviceType(rest...);
109 }
110}
111
112template <typename TensorLike, typename = void>
113class IsTensorLike : public std::false_type {};
114
115template <typename TensorLike>
116class IsTensorLike<
117 TensorLike,
118 std::void_t<decltype(std::declval<const TensorLike&>().data()),
119 decltype(std::declval<const TensorLike&>().shape()),
120 decltype(std::declval<const TensorLike&>().strides()),
121 decltype(std::declval<const TensorLike&>().dtype()),
122 decltype(std::declval<const TensorLike&>().device())>>
123 : public std::true_type {};
124
125template <typename T, typename std::enable_if_t<
126 IsTensorLike<std::decay_t<T>>::value, int> = 0>
127Tensor AsCallArg(const T& tensor) {
128 return Tensor{tensor};
129}
130
131template <typename T, typename std::enable_if_t<
132 !IsTensorLike<std::decay_t<T>>::value, int> = 0>
133const T& AsCallArg(const T& value) {
134 return value;
135}
136
137template <typename Key, typename TensorLike, typename Args, typename = void>
138class HasMakeReturnValueImpl : public std::false_type {};
139
140template <typename Key, typename TensorLike, typename... Args>
141class HasMakeReturnValueImpl<
142 Key, TensorLike, std::tuple<Args...>,
143 std::void_t<decltype(Key::MakeReturnValue(std::declval<const TensorLike&>(),
144 std::declval<const Args&>()...))>>
145 : public std::true_type {};
146
147template <typename Key, typename... Args>
148class HasMakeReturnValueImpl<Key, Tensor, std::tuple<Args...>>
149 : public std::false_type {};
150
151template <typename Key, typename TensorLike, typename... Args>
152class HasMakeReturnValue
153 : public HasMakeReturnValueImpl<Key, std::decay_t<TensorLike>,
154 std::tuple<Args...>> {};
155
157 static const bool enabled = [] {
158 const char* value = std::getenv("INFINI_OPS_TRACE_CALLS");
159 return value != nullptr && value[0] != '\0' && value[0] != '0';
160 }();
161 return enabled;
162}
163
164template <typename Key>
165constexpr std::string_view OperatorName() {
166#if defined(__clang__) || defined(__GNUC__)
167 std::string_view name{__PRETTY_FUNCTION__};
168 constexpr std::string_view marker{"Key = "};
169 const auto start = name.find(marker) + marker.size();
170 const auto end = name.find_first_of(";]", start);
171 name = name.substr(start, end - start);
172#else
173 std::string_view name{"unknown"};
174#endif
175 constexpr std::string_view namespace_prefix{"infini::ops::"};
176 if (name.rfind(namespace_prefix, 0) == 0) {
177 name.remove_prefix(namespace_prefix.size());
178 }
179 return name;
180}
181
182template <typename Key>
183void TraceOperatorCall(const CacheKey& key, const Config& config) {
184 if (!TraceOperatorCallsEnabled()) return;
185
186 const auto device =
187 key.tensors.empty()
188 ? std::string_view{"unknown"}
189 : Device::StringFromType(key.tensors.front().device().type());
190 constexpr auto operator_name = OperatorName<Key>();
191 std::fprintf(stderr,
192 "[INFINI_OPS_TRACE_CALLS] {\"operator_name\": \"%.*s\", "
193 "\"device_type\": \"%.*s\", \"implementation\": %zu}\n",
194 static_cast<int>(operator_name.size()), operator_name.data(),
195 static_cast<int>(device.size()), device.data(),
196 config.implementation_index());
197}
198
199} // namespace infini::ops::detail
200
201template <>
202struct std::hash<infini::ops::detail::CacheKey> {
203 std::size_t operator()(const infini::ops::detail::CacheKey& key) const {
204 return key.hash;
205 }
206};
207
208template <>
209struct std::equal_to<infini::ops::detail::CacheKey> {
210 bool operator()(const infini::ops::detail::CacheKey& a,
211 const infini::ops::detail::CacheKey& b) const {
212 if (a.scalar_hash != b.scalar_hash) return false;
213 if (a.tensors.size() != b.tensors.size()) return false;
214 std::equal_to<infini::ops::Tensor> eq;
215 for (std::size_t i = 0; i < a.tensors.size(); ++i) {
216 if (!eq(a.tensors[i], b.tensors[i])) return false;
217 }
218 return true;
219 }
220};
221
222namespace infini::ops {
223
224template <typename Key>
226 template <typename... Args>
227 detail::CacheKey operator()(const Config& config, const Args&... args) const {
228 return detail::CacheKey::Build(config.implementation_index(), args...);
229 }
230};
231
232namespace detail {
233
234template <typename Key, typename... Args>
235std::size_t ResolveImplementationIndex(const Config& config,
236 Device::Type dev_type,
237 const Args&... args);
238
239template <typename Key, typename... Args>
240std::size_t ResolveImplementationIndexOnline(const Handle& handle,
241 const Config& config,
242 Device::Type dev_type,
243 const Args&... args);
244
245} // namespace detail
246
247template <typename Key, Device::Type kDev>
248struct ActiveImplementations;
249
251 public:
252 virtual ~OperatorBase() = default;
253
254 virtual std::size_t workspace_size_in_bytes() const { return 0; }
255
256 void set_handle(const Handle& handle) { handle_ptr_ = handle.Clone(); }
257
258 void set_config(const Config& config) { config_ptr_ = config.Clone(); }
259
260 void set_stream(void* stream) { stream_ = stream; }
261
262 void set_workspace(void* workspace) { workspace_ = workspace; }
263
267
268 protected:
269 std::unique_ptr<Handle> handle_ptr_;
270
271 std::unique_ptr<Config> config_ptr_;
272
273 void* stream_{nullptr};
274
275 void* workspace_{nullptr};
276
278};
279
280template <typename Key, Device::Type device_type = Device::Type::kCount,
281 std::size_t implementation_index = 0>
282class Operator : public OperatorBase {
283 public:
284 // Invalidate the operator cache. Cached operators are destroyed on the next
285 // `Call()` invocation. Intended for test isolation; production code should
286 // never call this.
287 static void clear_cache() {
288 cache_generation_.fetch_add(1, std::memory_order_relaxed);
289 }
290
291 template <typename... Args>
292 static std::unique_ptr<Operator> Make(const Config& config,
293 const Tensor tensor, Args&&... args) {
294 const auto dev_type = tensor.device().type();
295 if (!TuningManager::Instance().IsEnabled() ||
297 return MakeWithDevice(config, dev_type, tensor,
298 std::forward<Args>(args)...);
299 }
300
301 auto resolved_config = config.Clone();
302 resolved_config->set_implementation_index(
303 detail::ResolveImplementationIndex<Key>(config, dev_type, tensor,
304 args...));
305 return MakeWithDevice(*resolved_config, dev_type, tensor,
306 std::forward<Args>(args)...);
307 }
308
309 template <typename... Args>
310 static std::unique_ptr<Operator> Make(const Tensor tensor, Args&&... args) {
311 return Make(ImplicitConfig(tensor.device().type()), tensor,
312 std::as_const(args)...);
313 }
314
315 template <typename... Args>
316 static std::unique_ptr<Operator> Make(const Config& config,
317 const std::vector<Tensor> tensors,
318 Args&&... args) {
319 assert(!tensors.empty() && "operator tensor list input cannot be empty");
320
321 const auto dev_type = tensors.front().device().type();
322 if (!TuningManager::Instance().IsEnabled() ||
324 return MakeWithDevice(config, dev_type, tensors,
325 std::forward<Args>(args)...);
326 }
327
328 auto resolved_config = config.Clone();
329 resolved_config->set_implementation_index(
330 detail::ResolveImplementationIndex<Key>(config, dev_type, tensors,
331 args...));
332 return MakeWithDevice(*resolved_config, dev_type, tensors,
333 std::forward<Args>(args)...);
334 }
335
336 template <typename... Args>
337 static std::unique_ptr<Operator> Make(const std::vector<Tensor> tensors,
338 Args&&... args) {
339 assert(!tensors.empty() && "operator tensor list input cannot be empty");
340
341 return Make(ImplicitConfig(tensors.front().device().type()), tensors,
342 std::as_const(args)...);
343 }
344
345 template <typename... Args>
346 static void Call(const Handle& handle, const Config& config,
347 const Args&... args);
348
349 template <typename... Args>
350 static void Call(const Tensor tensor, const Args&... args) {
351 return Call({}, ImplicitConfig(tensor.device().type()), tensor, args...);
352 }
353
354 template <
355 typename TensorLike, typename... Args,
356 typename std::enable_if_t<
357 detail::HasMakeReturnValue<Key, TensorLike, Args...>::value, int> = 0>
358 static auto Call(const TensorLike& tensor, const Args&... args) {
359 return CallReturning(tensor, args...);
360 }
361
362 static std::vector<std::size_t> active_implementation_indices(
363 Device::Type dev_type) {
364 if (!detail::ListContains(dev_type, ActiveDevices<Key>{})) {
365 return {};
366 }
367
368 std::vector<std::size_t> result;
369 DispatchFunc<ActiveDevices<Key>>(
370 dev_type,
371 [&](auto device_tag) {
372 constexpr Device::Type kDev = decltype(device_tag)::value;
373 result = detail::ListToVector(
375 },
376 "Operator::active_implementation_indices");
377 return result;
378 }
379
380 template <typename... Args>
381 void operator()(const Handle& handle, const Args&... args) {
382 set_handle(handle);
383 set_stream(handle.stream());
384 set_workspace(handle.workspace());
386
387 return operator()(args...);
388 }
389
390 template <typename... Args>
391 void operator()(const Args&... args) const {
392 return (*static_cast<const Key*>(this))(args...);
393 }
394
395 protected:
396 static constexpr Device::Type device_type_{device_type};
397
398 static constexpr std::size_t implementation_index_{implementation_index};
399
400 private:
401 template <auto first, auto... rest>
402 static constexpr std::size_t FirstActiveImplementationIndex(
403 List<first, rest...>) {
404 return static_cast<std::size_t>(first);
405 }
406
407 static std::size_t FirstActiveImplementationIndex(List<>) {
408 assert(false && "operator has no active implementation for this device");
409 std::abort();
410 }
411
412 static std::size_t DefaultImplementationIndex(Device::Type dev_type) {
413 std::size_t default_index{0};
414
415 DispatchFunc<ActiveDevices<Key>>(
416 dev_type,
417 [&](auto device_tag) {
418 constexpr Device::Type kDev = decltype(device_tag)::value;
419 default_index = FirstActiveImplementationIndex(
421 },
422 "Operator::DefaultImplementationIndex");
423
424 return default_index;
425 }
426
427 static Config DefaultConfig(Device::Type dev_type) {
428 Config config;
429 config.set_implementation_index(DefaultImplementationIndex(dev_type));
430
431 return config;
432 }
433
434 static Config ImplicitConfig(Device::Type dev_type) {
435 if (TuningManager::Instance().IsEnabled()) return Config{};
436 return DefaultConfig(dev_type);
437 }
438
439 template <typename TensorLike, typename... Args>
440 static auto CallReturning(const TensorLike& tensor, const Args&... args) {
441 auto out = Key::MakeReturnValue(tensor, args...);
442 Key::Call(detail::AsCallArg(tensor), detail::AsCallArg(args)...,
443 detail::AsCallArg(out));
444 return out;
445 }
446
447 template <typename... Args>
448 static std::unique_ptr<Operator> MakeWithDevice(
449 const Config& config, Device::Type dispatch_device_type, Args&&... args) {
450 std::unique_ptr<Operator> op_ptr;
451 auto cache_args = std::forward_as_tuple(args...);
452
453 DispatchFunc<ActiveDevices<Key>>(
454 dispatch_device_type,
455 [&](auto device_tag) {
456 constexpr Device::Type kDev = decltype(device_tag)::value;
458 config.implementation_index(),
459 [&](auto implementation_tag) {
460 constexpr std::size_t kImplementationIndex =
461 decltype(implementation_tag)::value;
462 if constexpr (std::is_constructible_v<
463 Operator<Key, kDev, kImplementationIndex>,
464 Args...>) {
465 std::apply(
466 [&](auto&... cached_args) {
467 op_ptr = std::make_unique<
468 Operator<Key, kDev, kImplementationIndex>>(
469 cached_args...);
470 },
471 cache_args);
472 } else {
473 assert(false &&
474 "operator is not implemented for this device and "
475 "implementation index");
476 }
477 },
478 "Operator::Make(implementation_index)",
480 },
481 "Operator::Make");
482
483 op_ptr->set_config(config);
484
485 return op_ptr;
486 }
487
488 static inline std::atomic<std::size_t> cache_generation_{0};
489};
490
491// Maximum number of implementation slots per (operator, device) pair.
492// Increase this value when adding operators with more implementations.
493constexpr std::size_t kMaxImplementations = 32;
494
495// SFINAE-based implementation detection. A partial specialization
496// `Operator<Key, kDev, N>` inherits from `Key` (the operator base class),
497// while the unspecialized primary template inherits only from `OperatorBase`.
498// `std::is_base_of` distinguishes the two at compile time, eliminating the
499// need for manual `registry.h` files.
500template <typename Key, Device::Type kDev, std::size_t N,
501 bool = std::is_base_of_v<Key, Operator<Key, kDev, N>>>
503 using type = List<>;
504};
505
506template <typename Key, Device::Type kDev, std::size_t N>
507struct ActiveImplementationsImpl<Key, kDev, N, true> {
508 using type = List<N>;
509};
510
511namespace detail {
512
513template <typename Key, Device::Type kDev, typename Seq>
514struct ActiveImplementationsHelper;
515
516template <typename Key, Device::Type kDev, std::size_t... ns>
517struct ActiveImplementationsHelper<Key, kDev, std::index_sequence<ns...>> {
518 using type = typename Flatten<
520};
521
522} // namespace detail
523
524template <typename Key, Device::Type kDev>
525struct ActiveImplementations {
526 using type = typename detail::ActiveImplementationsHelper<
527 Key, kDev, std::make_index_sequence<kMaxImplementations>>::type;
528};
529
530namespace detail {
531
532template <typename Key, typename... Args>
533std::size_t ResolveImplementationIndex(const Config& config,
534 Device::Type dev_type,
535 const Args&... args) {
536 if (!config.needs_implementation_resolution()) {
537 return config.implementation_index();
538 }
539
540 auto indices = Operator<Key>::active_implementation_indices(dev_type);
541 if (indices.empty()) return config.implementation_index();
542
543 auto signature = TuningSignature::Build(args...);
544 constexpr auto op_name = OperatorName<Key>();
545 auto tuned_index =
546 TuningManager::Instance().Lookup(op_name, dev_type, signature);
547 auto chosen = indices.front();
548
549 if (tuned_index.has_value()) {
550 bool is_valid = std::find(indices.begin(), indices.end(), *tuned_index) !=
551 indices.end();
552 if (is_valid) {
553 chosen = *tuned_index;
554 } else {
555 std::cerr << "[Tuning] Warning: tuned implementation " << *tuned_index
556 << " for " << op_name << " on "
557 << Device::StringFromType(dev_type)
558 << " is not available (compiled indices:";
559 for (auto idx : indices) std::cerr << " " << idx;
560 std::cerr << "), falling back to " << chosen << std::endl;
561 }
562 }
563
564 return chosen;
565}
566
567template <typename Key, typename... Args>
568double BenchmarkImplementation(const Handle& handle, Device::Type dev_type,
569 std::size_t impl_index, const Args&... args) {
570 Config fixed;
571 fixed.set_implementation_index(impl_index);
572
573 auto op = Operator<Key>::Make(fixed, args...);
574
575 const auto& tuning = TuningManager::Instance();
576 const int warmup = tuning.warmup_count();
577 const int repeat = tuning.repeat_count();
578
579 for (int i = 0; i < warmup; ++i) {
580 (*op)(handle, args...);
581 }
582 SyncDevice(dev_type);
583
584 double best = std::numeric_limits<double>::infinity();
585 for (int i = 0; i < repeat; ++i) {
586 auto start = std::chrono::steady_clock::now();
587 (*op)(handle, args...);
588 SyncDevice(dev_type);
589 auto end = std::chrono::steady_clock::now();
590 double elapsed = std::chrono::duration<double>(end - start).count();
591 best = std::min(best, elapsed);
592 }
593 return best;
594}
595
596template <typename Key, typename... Args>
597std::size_t ResolveImplementationIndexOnline(const Handle& handle,
598 const Config& config,
599 Device::Type dev_type,
600 const Args&... args) {
601 if (!config.needs_implementation_resolution()) {
602 return config.implementation_index();
603 }
604
605 auto& tuning = TuningManager::Instance();
606 if (!tuning.IsEnabled()) {
607 return ResolveImplementationIndex<Key>(config, dev_type, args...);
608 }
609
610 auto indices = Operator<Key>::active_implementation_indices(dev_type);
611 if (indices.empty()) return config.implementation_index();
612
613 auto signature = TuningSignature::Build(args...);
614 constexpr auto op_name = OperatorName<Key>();
615 auto tuned = tuning.Lookup(op_name, dev_type, signature);
616 std::size_t chosen;
617
618 if (tuned.has_value() &&
619 std::find(indices.begin(), indices.end(), *tuned) != indices.end()) {
620 chosen = *tuned;
621 } else if (indices.size() == 1) {
622 chosen = indices.front();
623 tuning.Record(op_name, dev_type, signature, chosen);
624 std::cout << "[Tuning] " << op_name << " on "
625 << Device::StringFromType(dev_type)
626 << ": single impl, chose index " << chosen << std::endl;
627 } else {
628 chosen = indices.front();
629 double best_time = std::numeric_limits<double>::infinity();
630 for (auto idx : indices) {
631 double time =
632 BenchmarkImplementation<Key>(handle, dev_type, idx, args...);
633 if (time < best_time) {
634 best_time = time;
635 chosen = idx;
636 }
637 }
638 tuning.Record(op_name, dev_type, signature, chosen);
639 std::cout << "[Tuning] " << op_name << " on "
640 << Device::StringFromType(dev_type) << ": benchmarked "
641 << indices.size() << " impls, chose index " << chosen << " ("
642 << best_time * 1e6 << " us)" << std::endl;
643 }
644
645 return chosen;
646}
647
648} // namespace detail
649
650} // namespace infini::ops
651
652#endif
Definition generated/include/config.h:12
void set_implementation_index(std::size_t implementation_index)
Definition generated/include/config.h:24
virtual std::unique_ptr< Config > Clone() const
Definition generated/include/config.h:16
bool needs_implementation_resolution() const
Definition generated/include/config.h:28
std::size_t implementation_index() const
Definition generated/include/config.h:20
Definition handle.h:11
void * stream() const
Definition handle.h:19
void * workspace() const
Definition handle.h:21
std::size_t workspace_size_in_bytes() const
Definition handle.h:23
virtual std::unique_ptr< Handle > Clone() const
Definition handle.h:15
Definition generated/include/operator.h:250
void set_workspace(void *workspace)
Definition generated/include/operator.h:262
void set_config(const Config &config)
Definition generated/include/operator.h:258
void * workspace_
Definition generated/include/operator.h:275
virtual std::size_t workspace_size_in_bytes() const
Definition generated/include/operator.h:254
std::size_t workspace_size_in_bytes_
Definition generated/include/operator.h:277
std::unique_ptr< Config > config_ptr_
Definition generated/include/operator.h:271
std::unique_ptr< Handle > handle_ptr_
Definition generated/include/operator.h:269
void set_handle(const Handle &handle)
Definition generated/include/operator.h:256
virtual ~OperatorBase()=default
void set_workspace_size_in_bytes(std::size_t workspace_size_in_bytes)
Definition generated/include/operator.h:264
void set_stream(void *stream)
Definition generated/include/operator.h:260
void * stream_
Definition generated/include/operator.h:273
Definition generated/include/operator.h:282
static void Call(const Tensor tensor, const Args &... args)
Definition generated/include/operator.h:350
static void clear_cache()
Definition generated/include/operator.h:287
static void Call(const Handle &handle, const Config &config, const Args &... args)
static std::unique_ptr< Operator > Make(const std::vector< Tensor > tensors, Args &&... args)
Definition generated/include/operator.h:337
void operator()(const Args &... args) const
Definition generated/include/operator.h:391
static std::unique_ptr< Operator > Make(const Config &config, const std::vector< Tensor > tensors, Args &&... args)
Definition generated/include/operator.h:316
static std::unique_ptr< Operator > Make(const Tensor tensor, Args &&... args)
Definition generated/include/operator.h:310
void operator()(const Handle &handle, const Args &... args)
Definition generated/include/operator.h:381
static std::unique_ptr< Operator > Make(const Config &config, const Tensor tensor, Args &&... args)
Definition generated/include/operator.h:292
static auto Call(const TensorLike &tensor, const Args &... args)
Definition generated/include/operator.h:358
static constexpr std::size_t implementation_index_
Definition generated/include/operator.h:398
static std::vector< std::size_t > active_implementation_indices(Device::Type dev_type)
Definition generated/include/operator.h:362
static constexpr Device::Type device_type_
Definition generated/include/operator.h:396
Definition generated/include/operator.h:28
void TraceOperatorCall(const CacheKey &key, const Config &config)
Definition generated/include/operator.h:183
std::size_t ResolveImplementationIndexOnline(const Handle &handle, const Config &config, Device::Type dev_type, const Args &... args)
Definition generated/include/operator.h:597
bool TraceOperatorCallsEnabled()
Definition generated/include/operator.h:156
std::size_t ResolveImplementationIndex(const Config &config, Device::Type dev_type, const Args &... args)
Definition generated/include/operator.h:533
double BenchmarkImplementation(const Handle &handle, Device::Type dev_type, std::size_t impl_index, const Args &... args)
Definition generated/include/operator.h:568
Device::Type FirstDeviceType()
Definition generated/include/operator.h:97
std::vector< std::size_t > ListToVector(List< values... >)
Definition generated/include/operator.h:78
Tensor AsCallArg(const T &tensor)
Definition generated/include/operator.h:127
auto DispatchImplementation(std::size_t implementation_index, Functor &&func, std::string_view context_str, List< implementation_indices... >, Args &&... args)
Definition generated/include/operator.h:68
constexpr std::string_view OperatorName()
Definition generated/include/operator.h:165
bool ListContains(ValueType value, List< values... >)
Definition generated/include/operator.h:83
void SyncDevice(Device::Type dev_type)
Definition generated/include/operator.h:87
Definition generated/include/operator.h:28
typename ActiveDevicesImpl< T >::type ActiveDevices
Definition device.h:33
constexpr std::size_t kMaxImplementations
Definition generated/include/operator.h:493
infini::rt::TensorView Tensor
Definition tensor.h:8
auto DispatchFunc(ValueType value, Functor &&func, std::string_view context_str="", Args &&... args)
Definition dispatcher.h:112
Definition generated/include/operator.h:28
List< N > type
Definition generated/include/operator.h:508
Definition generated/include/operator.h:502
List<> type
Definition generated/include/operator.h:503
typename detail::ActiveImplementationsHelper< Key, kDev, std::make_index_sequence< kMaxImplementations > >::type type
Definition generated/include/operator.h:527
Definition generated/include/operator.h:225
detail::CacheKey operator()(const Config &config, const Args &... args) const
Definition generated/include/operator.h:227