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