1#ifndef INFINI_OPS_OPERATOR_H_
2#define INFINI_OPS_OPERATOR_H_
17#include <unordered_map>
33 std::vector<Tensor> tensors;
35 std::size_t scalar_hash;
37 template <
typename... Args>
38 static CacheKey Build(
const Args&... args) {
42 (key.Absorb(args), ...);
47 void Absorb(
const Tensor& t) {
52 void Absorb(
const std::vector<Tensor>& ts) {
53 HashCombine(hash, ts.size());
54 for (
const auto& t : ts) {
61 void Absorb(
const T& v) {
63 HashCombine(scalar_hash, v);
67template <
typename Functor,
typename... Args,
auto... implementation_indices>
69 std::string_view context_str,
70 List<implementation_indices...>, Args&&... args) {
72 static_cast<std::size_t
>(implementation_indices)...>(
73 implementation_index, std::forward<Functor>(func), context_str,
74 std::forward<Args>(args)...);
77template <
auto... values>
79 return {
static_cast<std::size_t
>(values)...};
82template <
typename ValueType,
auto... values>
84 return ((value ==
static_cast<ValueType
>(values)) || ...);
88 DispatchFunc<ActiveDevices<void>>(
91 constexpr Device::Type kDev =
decltype(device_tag)::value;
92 infini::rt::runtime::Runtime<kDev>::DeviceSynchronize();
99template <
typename First,
typename... 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>>) {
106 : first.front().device().type();
112template <
typename TensorLike,
typename =
void>
113class IsTensorLike :
public std::false_type {};
115template <
typename 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 {};
125template <
typename T,
typename std::enable_if_t<
126 IsTensorLike<std::decay_t<T>>::value,
int> = 0>
131template <
typename T,
typename std::enable_if_t<
132 !IsTensorLike<std::decay_t<T>>::value,
int> = 0>
137template <
typename Key,
typename TensorLike,
typename Args,
typename =
void>
138class HasMakeReturnValueImpl :
public std::false_type {};
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 {};
147template <
typename Key,
typename... Args>
148class HasMakeReturnValueImpl<Key,
Tensor, std::tuple<Args...>>
149 :
public std::false_type {};
151template <
typename Key,
typename TensorLike,
typename... Args>
152class HasMakeReturnValue
153 :
public HasMakeReturnValueImpl<Key, std::decay_t<TensorLike>,
154 std::tuple<Args...>> {};
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';
164template <
typename Key>
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);
173 std::string_view name{
"unknown"};
175 constexpr std::string_view namespace_prefix{
"infini::ops::"};
176 if (name.rfind(namespace_prefix, 0) == 0) {
177 name.remove_prefix(namespace_prefix.size());
182template <
typename Key>
188 ? std::string_view{
"unknown"}
189 : Device::StringFromType(key.tensors.front().device().type());
190 constexpr auto operator_name = OperatorName<Key>();
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(),
202struct std::hash<
infini::ops::detail::CacheKey> {
203 std::size_t operator()(
const infini::ops::detail::CacheKey& key)
const {
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;
224template <
typename Key>
226 template <
typename... Args>
234template <
typename Key,
typename... Args>
236 Device::Type dev_type,
237 const Args&... args);
239template <
typename Key,
typename... Args>
241 const Config& config,
242 Device::Type dev_type,
243 const Args&... args);
247template <
typename Key, Device::Type kDev>
248struct ActiveImplementations;
280template <
typename Key, Device::Type device_type = Device::Type::kCount,
281 std::size_t implementation_index = 0>
288 cache_generation_.fetch_add(1, std::memory_order_relaxed);
291 template <
typename... Args>
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)...);
301 auto resolved_config = config.
Clone();
302 resolved_config->set_implementation_index(
303 detail::ResolveImplementationIndex<Key>(config, dev_type, tensor,
305 return MakeWithDevice(*resolved_config, dev_type, tensor,
306 std::forward<Args>(args)...);
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)...);
315 template <
typename... Args>
317 const std::vector<Tensor> tensors,
319 assert(!tensors.empty() &&
"operator tensor list input cannot be empty");
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)...);
328 auto resolved_config = config.
Clone();
329 resolved_config->set_implementation_index(
330 detail::ResolveImplementationIndex<Key>(config, dev_type, tensors,
332 return MakeWithDevice(*resolved_config, dev_type, tensors,
333 std::forward<Args>(args)...);
336 template <
typename... Args>
337 static std::unique_ptr<Operator>
Make(
const std::vector<Tensor> tensors,
339 assert(!tensors.empty() &&
"operator tensor list input cannot be empty");
341 return Make(ImplicitConfig(tensors.front().device().type()), tensors,
342 std::as_const(args)...);
345 template <
typename... Args>
347 const Args&... args);
349 template <
typename... Args>
351 return Call({}, ImplicitConfig(tensor.device().type()), tensor, args...);
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...);
363 Device::Type dev_type) {
368 std::vector<std::size_t> result;
369 DispatchFunc<ActiveDevices<Key>>(
371 [&](
auto device_tag) {
372 constexpr Device::Type kDev =
decltype(device_tag)::value;
376 "Operator::active_implementation_indices");
380 template <
typename... Args>
390 template <
typename... Args>
392 return (*
static_cast<const Key*
>(
this))(args...);
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);
407 static std::size_t FirstActiveImplementationIndex(List<>) {
408 assert(
false &&
"operator has no active implementation for this device");
412 static std::size_t DefaultImplementationIndex(Device::Type dev_type) {
413 std::size_t default_index{0};
415 DispatchFunc<ActiveDevices<Key>>(
417 [&](
auto device_tag) {
418 constexpr Device::Type kDev =
decltype(device_tag)::value;
419 default_index = FirstActiveImplementationIndex(
422 "Operator::DefaultImplementationIndex");
424 return default_index;
427 static Config DefaultConfig(Device::Type dev_type) {
429 config.set_implementation_index(DefaultImplementationIndex(dev_type));
434 static Config ImplicitConfig(Device::Type dev_type) {
435 if (TuningManager::Instance().IsEnabled())
return Config{};
436 return DefaultConfig(dev_type);
439 template <
typename TensorLike,
typename... Args>
440 static auto CallReturning(
const TensorLike& tensor,
const Args&... args) {
441 auto out = Key::MakeReturnValue(tensor, args...);
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...);
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>,
466 [&](auto&... cached_args) {
467 op_ptr = std::make_unique<
468 Operator<Key, kDev, kImplementationIndex>>(
474 "operator is not implemented for this device and "
475 "implementation index");
478 "Operator::Make(implementation_index)",
483 op_ptr->set_config(config);
488 static inline std::atomic<std::size_t> cache_generation_{0};
500template <
typename Key, Device::Type kDev, std::size_t N,
501 bool = std::is_base_of_v<Key, Operator<Key, kDev, N>>>
506template <
typename Key, Device::Type kDev, std::
size_t N>
513template <
typename Key, Device::Type kDev,
typename Seq>
514struct ActiveImplementationsHelper;
516template <
typename Key, Device::Type kDev, std::size_t... ns>
517struct ActiveImplementationsHelper<Key, kDev, std::index_sequence<ns...>> {
518 using type =
typename Flatten<
524template <
typename Key, Device::Type kDev>
525struct ActiveImplementations {
526 using type =
typename detail::ActiveImplementationsHelper<
527 Key, kDev, std::make_index_sequence<kMaxImplementations>>::type;
532template <
typename Key,
typename... Args>
534 Device::Type dev_type,
535 const Args&... args) {
543 auto signature = TuningSignature::Build(args...);
544 constexpr auto op_name = OperatorName<Key>();
546 TuningManager::Instance().Lookup(op_name, dev_type, signature);
547 auto chosen = indices.front();
549 if (tuned_index.has_value()) {
550 bool is_valid = std::find(indices.begin(), indices.end(), *tuned_index) !=
553 chosen = *tuned_index;
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;
567template <
typename Key,
typename... Args>
569 std::size_t impl_index,
const Args&... args) {
575 const auto& tuning = TuningManager::Instance();
576 const int warmup = tuning.warmup_count();
577 const int repeat = tuning.repeat_count();
579 for (
int i = 0; i < warmup; ++i) {
580 (*op)(handle, args...);
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...);
589 auto end = std::chrono::steady_clock::now();
590 double elapsed = std::chrono::duration<double>(end - start).count();
591 best = std::min(best, elapsed);
596template <
typename Key,
typename... Args>
599 Device::Type dev_type,
600 const Args&... args) {
605 auto& tuning = TuningManager::Instance();
606 if (!tuning.IsEnabled()) {
607 return ResolveImplementationIndex<Key>(config, dev_type, args...);
613 auto signature = TuningSignature::Build(args...);
614 constexpr auto op_name = OperatorName<Key>();
615 auto tuned = tuning.Lookup(op_name, dev_type, signature);
618 if (tuned.has_value() &&
619 std::find(indices.begin(), indices.end(), *tuned) != indices.end()) {
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;
628 chosen = indices.front();
629 double best_time = std::numeric_limits<double>::infinity();
630 for (
auto idx : indices) {
632 BenchmarkImplementation<Key>(handle, dev_type, idx, args...);
633 if (time < best_time) {
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;
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
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