1#ifndef INFINI_OPS_OPERATOR_H_
2#define INFINI_OPS_OPERATOR_H_
17#include <unordered_map>
24#include "host_range_profiler.h"
34 std::vector<Tensor> tensors;
36 std::size_t scalar_hash;
38 template <
typename... Args>
39 static CacheKey Build(
const Args&... args) {
43 (key.Absorb(args), ...);
48 void Absorb(
const Tensor& t) {
53 void Absorb(
const std::vector<Tensor>& ts) {
54 HashCombine(hash, ts.size());
55 for (
const auto& t : ts) {
62 void Absorb(
const T& v) {
64 HashCombine(scalar_hash, v);
68template <
typename Functor,
typename... Args,
auto... implementation_indices>
70 std::string_view context_str,
71 List<implementation_indices...>, Args&&... args) {
73 static_cast<std::size_t
>(implementation_indices)...>(
74 implementation_index, std::forward<Functor>(func), context_str,
75 std::forward<Args>(args)...);
78template <
auto... values>
80 return {
static_cast<std::size_t
>(values)...};
83template <
typename ValueType,
auto... values>
85 return ((value ==
static_cast<ValueType
>(values)) || ...);
89 DispatchFunc<ActiveDevices<void>>(
92 constexpr Device::Type kDev =
decltype(device_tag)::value;
93 infini::rt::runtime::Runtime<kDev>::DeviceSynchronize();
100template <
typename First,
typename... 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>>) {
107 : first.front().device().type();
113template <
typename TensorLike,
typename =
void>
114class IsTensorLike :
public std::false_type {};
116template <
typename 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 {};
126template <
typename T,
typename std::enable_if_t<
127 IsTensorLike<std::decay_t<T>>::value,
int> = 0>
132template <
typename T,
typename std::enable_if_t<
133 !IsTensorLike<std::decay_t<T>>::value,
int> = 0>
138template <
typename Key,
typename TensorLike,
typename Args,
typename =
void>
139class HasMakeReturnValueImpl :
public std::false_type {};
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 {};
148template <
typename Key,
typename... Args>
149class HasMakeReturnValueImpl<Key,
Tensor, std::tuple<Args...>>
150 :
public std::false_type {};
152template <
typename Key,
typename TensorLike,
typename... Args>
153class HasMakeReturnValue
154 :
public HasMakeReturnValueImpl<Key, std::decay_t<TensorLike>,
155 std::tuple<Args...>> {};
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';
165template <
typename Key>
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);
174 std::string_view name{
"unknown"};
176 constexpr std::string_view namespace_prefix{
"infini::ops::"};
177 if (name.rfind(namespace_prefix, 0) == 0) {
178 name.remove_prefix(namespace_prefix.size());
183template <
typename Key>
189 ? std::string_view{
"unknown"}
190 : Device::StringFromType(key.tensors.front().device().type());
191 constexpr auto operator_name = OperatorName<Key>();
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());
203struct std::hash<
infini::ops::detail::CacheKey> {
204 std::size_t operator()(
const infini::ops::detail::CacheKey& key)
const {
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;
225template <
typename Key>
226struct CacheKeyBuilder {
227 template <
typename... Args>
235template <
typename Key,
typename... Args>
237 Device::Type dev_type,
238 const Args&... args);
240template <
typename Key,
typename... Args>
242 const Config& config,
243 Device::Type dev_type,
244 const Args&... args);
248template <
typename Key, Device::Type kDev>
249struct ActiveImplementations;
281template <
typename Key, Device::Type device_type = Device::Type::kCount,
282 std::size_t implementation_index = 0>
283class Operator :
public OperatorBase {
289 cache_generation_.fetch_add(1, std::memory_order_relaxed);
292 template <
typename... Args>
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)...);
302 auto resolved_config = config.
Clone();
303 resolved_config->set_implementation_index(
304 detail::ResolveImplementationIndex<Key>(config, dev_type, tensor,
306 return MakeWithDevice(*resolved_config, dev_type, tensor,
307 std::forward<Args>(args)...);
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)...);
316 template <
typename... Args>
318 const std::vector<Tensor> tensors,
320 assert(!tensors.empty() &&
"operator tensor list input cannot be empty");
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)...);
329 auto resolved_config = config.
Clone();
330 resolved_config->set_implementation_index(
331 detail::ResolveImplementationIndex<Key>(config, dev_type, tensors,
333 return MakeWithDevice(*resolved_config, dev_type, tensors,
334 std::forward<Args>(args)...);
337 template <
typename... Args>
338 static std::unique_ptr<Operator>
Make(
const std::vector<Tensor> tensors,
340 assert(!tensors.empty() &&
"operator tensor list input cannot be empty");
342 return Make(ImplicitConfig(tensors.front().device().type()), tensors,
343 std::as_const(args)...);
346 template <
typename... Args>
348 const Args&... args) {
349 [[maybe_unused]] HostRangeScope host_range_operator_call{
350 HostRangeLayer::kOperatorCall};
352 static thread_local std::unordered_map<detail::CacheKey,
353 std::unique_ptr<Operator>>
355 static thread_local std::size_t generation{0};
357 const auto cache_generation =
358 cache_generation_.load(std::memory_order_relaxed);
359 if (generation != cache_generation) {
361 generation = cache_generation;
364 std::unique_ptr<Config> resolved_config;
365 const Config* effective_config = &config;
366 if (TuningManager::Instance().IsEnabled() &&
369 assert(dev_type != Device::Type::kCount &&
370 "operator call requires at least one tensor argument");
372 const auto resolved_implementation_index =
373 detail::ResolveImplementationIndexOnline<Key>(handle, config,
375 resolved_config = config.
Clone();
376 resolved_config->set_implementation_index(resolved_implementation_index);
377 effective_config = resolved_config.get();
380#if defined(INFINI_OPS_ENABLE_HOST_RANGE_PROFILING)
382 HostRangeScope host_range_cache_key{HostRangeLayer::kCacheKey};
385 detail::TraceOperatorCall<Key>(key, *effective_config);
388 HostRangeScope host_range_cache_lookup{HostRangeLayer::kCacheLookup};
389 return cache.find(key);
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;
400 detail::TraceOperatorCall<Key>(key, *effective_config);
402 auto it{cache.find(key)};
404 if (it == cache.end()) {
406 cache.emplace(std::move(key),
Make(*effective_config, args...)).first;
410 auto& op{it->second};
412 [[maybe_unused]] HostRangeScope host_range_operator_invoke{
413 HostRangeLayer::kOperatorInvoke};
414 return (*op)(handle, args...);
417 template <
typename... Args>
419 return Call({}, ImplicitConfig(tensor.device().type()), tensor, args...);
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...);
431 Device::Type dev_type) {
436 std::vector<std::size_t> result;
437 DispatchFunc<ActiveDevices<Key>>(
439 [&](
auto device_tag) {
440 constexpr Device::Type kDev =
decltype(device_tag)::value;
444 "Operator::active_implementation_indices");
448 template <
typename... Args>
458 template <
typename... Args>
460 return (*
static_cast<const Key*
>(
this))(args...);
464 static constexpr Device::Type
device_type_{device_type};
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);
475 static std::size_t FirstActiveImplementationIndex(List<>) {
476 assert(
false &&
"operator has no active implementation for this device");
480 static std::size_t DefaultImplementationIndex(Device::Type dev_type) {
481 std::size_t default_index{0};
483 DispatchFunc<ActiveDevices<Key>>(
485 [&](
auto device_tag) {
486 constexpr Device::Type kDev =
decltype(device_tag)::value;
487 default_index = FirstActiveImplementationIndex(
490 "Operator::DefaultImplementationIndex");
492 return default_index;
495 static Config DefaultConfig(Device::Type dev_type) {
497 config.set_implementation_index(DefaultImplementationIndex(dev_type));
502 static Config ImplicitConfig(Device::Type dev_type) {
503 if (TuningManager::Instance().IsEnabled())
return Config{};
504 return DefaultConfig(dev_type);
507 template <
typename TensorLike,
typename... Args>
508 static auto CallReturning(
const TensorLike& tensor,
const Args&... args) {
509 auto out = Key::MakeReturnValue(tensor, args...);
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...);
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>,
534 [&](auto&... cached_args) {
535 op_ptr = std::make_unique<
536 Operator<Key, kDev, kImplementationIndex>>(
542 "operator is not implemented for this device and "
543 "implementation index");
546 "Operator::Make(implementation_index)",
551 op_ptr->set_config(config);
556 static inline std::atomic<std::size_t> cache_generation_{0};
568template <
typename Key, Device::Type kDev, std::size_t N,
569 bool = std::is_base_of_v<Key, Operator<Key, kDev, N>>>
570struct ActiveImplementationsImpl {
574template <
typename Key, Device::Type kDev, std::
size_t N>
581template <
typename Key, Device::Type kDev,
typename Seq>
582struct ActiveImplementationsHelper;
584template <
typename Key, Device::Type kDev, std::size_t... ns>
585struct ActiveImplementationsHelper<Key, kDev, std::index_sequence<ns...>> {
586 using type =
typename Flatten<
592template <
typename Key, Device::Type kDev>
594 using type =
typename detail::ActiveImplementationsHelper<
595 Key, kDev, std::make_index_sequence<kMaxImplementations>>::type;
600template <
typename Key,
typename... Args>
601std::size_t ResolveImplementationIndex(
const Config& config,
602 Device::Type dev_type,
603 const Args&... args) {
611 auto signature = TuningSignature::Build(args...);
612 constexpr auto op_name = OperatorName<Key>();
614 TuningManager::Instance().Lookup(op_name, dev_type, signature);
615 auto chosen = indices.front();
617 if (tuned_index.has_value()) {
618 bool is_valid = std::find(indices.begin(), indices.end(), *tuned_index) !=
621 chosen = *tuned_index;
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;
635template <
typename Key,
typename... Args>
636double BenchmarkImplementation(
const Handle& handle, Device::Type dev_type,
637 std::size_t impl_index,
const Args&... args) {
639 fixed.set_implementation_index(impl_index);
641 auto op = Operator<Key>::Make(fixed, args...);
643 const auto& tuning = TuningManager::Instance();
644 const int warmup = tuning.warmup_count();
645 const int repeat = tuning.repeat_count();
647 for (
int i = 0; i < warmup; ++i) {
648 (*op)(handle, args...);
650 SyncDevice(dev_type);
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);
664template <
typename Key,
typename... Args>
666 const Config& config,
667 Device::Type dev_type,
668 const Args&... args) {
669 if (!config.needs_implementation_resolution()) {
670 return config.implementation_index();
673 auto& tuning = TuningManager::Instance();
674 if (!tuning.IsEnabled()) {
675 return ResolveImplementationIndex<Key>(config, dev_type, args...);
678 auto indices = Operator<Key>::active_implementation_indices(dev_type);
679 if (indices.empty())
return config.implementation_index();
681 auto signature = TuningSignature::Build(args...);
682 constexpr auto op_name = OperatorName<Key>();
683 auto tuned = tuning.Lookup(op_name, dev_type, signature);
686 if (tuned.has_value() &&
687 std::find(indices.begin(), indices.end(), *tuned) != indices.end()) {
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;
696 chosen = indices.front();
697 double best_time = std::numeric_limits<double>::infinity();
698 for (
auto idx : indices) {
700 BenchmarkImplementation<Key>(handle, dev_type, idx, args...);
701 if (time < best_time) {
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;
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
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