5 #ifndef GKO_PUBLIC_CORE_BASE_ABSTRACT_FACTORY_HPP_
6 #define GKO_PUBLIC_CORE_BASE_ABSTRACT_FACTORY_HPP_
10 #include <unordered_map>
12 #include <ginkgo/core/base/polymorphic_object.hpp>
34 using type =
typename T::element_type;
38 using element_type_t =
typename element_type<T>::type;
49 using value_element_type_t =
typename T::value_type::element_type;
76 template <
typename AbstractProductType,
typename ComponentsType>
79 using abstract_product_type = AbstractProductType;
80 using components_type = ComponentsType;
96 template <
typename... Args>
97 std::unique_ptr<abstract_product_type>
generate(Args&&... args)
const
100 this->generate_impl(components_type{std::forward<Args>(args)...});
101 for (
auto logger : this->loggers_) {
102 product->add_logger(logger);
124 virtual std::unique_ptr<abstract_product_type> generate_impl(
125 ComponentsType args)
const = 0;
152 template <
typename ConcreteFactory,
typename ProductType,
153 typename ParametersType,
typename PolymorphicBase>
156 using product_type = ProductType;
157 using parameters_type = ParametersType;
158 using polymorphic_base = PolymorphicBase;
159 using abstract_product_type =
160 typename PolymorphicBase::abstract_product_type;
161 using components_type =
typename PolymorphicBase::components_type;
163 template <
typename... Args>
164 std::unique_ptr<product_type> generate(Args&&... args)
const
166 auto product = std::unique_ptr<product_type>(static_cast<product_type*>(
167 this->polymorphic_base::generate(std::forward<Args>(args)...)
194 static parameters_type
create() {
return {}; }
204 const parameters_type& parameters = {})
205 : PolymorphicBase(std::move(exec)), parameters_{parameters}
208 std::unique_ptr<abstract_product_type> generate_impl(
209 components_type args)
const override
211 return std::unique_ptr<abstract_product_type>(
212 new product_type(
self(), args));
216 GKO_ENABLE_SELF(ConcreteFactory);
218 ParametersType parameters_;
234 template <
typename ConcreteParametersType,
typename Factory>
237 using factory = Factory;
243 template <
typename... Args>
246 this->loggers = {std::forward<Args>(_value)...};
257 std::unique_ptr<Factory>
on(std::shared_ptr<const Executor> exec)
const
259 ConcreteParametersType copy = *
self();
260 for (
const auto& item : deferred_factories) {
261 item.second(exec, copy);
263 auto factory = std::unique_ptr<Factory>(
new Factory(exec, copy));
264 for (
auto& logger : loggers) {
265 factory->add_logger(logger);
271 GKO_ENABLE_SELF(ConcreteParametersType);
276 std::vector<std::shared_ptr<const log::Logger>> loggers{};
284 std::unordered_map<std::string,
285 std::function<void(std::shared_ptr<const Executor> exec,
286 ConcreteParametersType&)>>
304 #define GKO_CREATE_FACTORY_PARAMETERS(_parameters_name, _factory_name) \
306 class _factory_name; \
307 struct _parameters_name##_type \
308 : public ::gko::enable_parameters_type<_parameters_name##_type, \
317 template <
typename From,
typename To>
318 struct is_pointer_convertible : std::is_convertible<From*, To*> {};
332 template <
typename FactoryType>
341 generator_ = [](std::shared_ptr<const Executor>) {
return nullptr; };
348 template <
typename ConcreteFactoryType,
349 std::enable_if_t<detail::is_pointer_convertible<
350 ConcreteFactoryType, FactoryType>::value>* =
nullptr>
353 generator_ = [factory =
354 std::shared_ptr<FactoryType>(std::move(factory))](
355 std::shared_ptr<const Executor>) {
return factory; };
362 template <
typename ConcreteFactoryType,
typename Deleter,
363 std::enable_if_t<detail::is_pointer_convertible<
364 ConcreteFactoryType, FactoryType>::value>* =
nullptr>
366 std::unique_ptr<ConcreteFactoryType, Deleter> factory)
368 generator_ = [factory =
369 std::shared_ptr<FactoryType>(std::move(factory))](
370 std::shared_ptr<const Executor>) {
return factory; };
378 template <
typename ParametersType,
379 typename U = decltype(std::declval<ParametersType>().
on(
380 std::shared_ptr<const Executor>{})),
381 std::enable_if_t<detail::is_pointer_convertible<
382 typename U::element_type, FactoryType>::value>* =
nullptr>
385 generator_ = [parameters](std::shared_ptr<const Executor> exec)
386 -> std::shared_ptr<FactoryType> {
return parameters.on(exec); };
393 std::shared_ptr<FactoryType>
on(std::shared_ptr<const Executor> exec)
const
396 GKO_NOT_SUPPORTED(*
this);
398 return generator_(exec);
402 bool is_empty()
const {
return !bool(generator_); }
405 std::function<std::shared_ptr<FactoryType>(std::shared_ptr<const Executor>)>
418 #define GKO_ENABLE_BUILD_METHOD(_factory_name) \
419 static auto build()->decltype(_factory_name::create()) \
421 return _factory_name::create(); \
423 static_assert(true, \
424 "This assert is used to counter the false positive extra " \
425 "semi-colon warnings")
428 #if !(defined(__CUDACC__) || defined(__HIPCC__))
441 #define GKO_FACTORY_PARAMETER(_name, ...) \
442 _name{__VA_ARGS__}; \
444 template <typename... Args> \
445 auto with_##_name(Args&&... _value) \
446 ->std::decay_t<decltype(*(this->self()))>& \
448 using type = decltype(this->_name); \
449 this->_name = type{std::forward<Args>(_value)...}; \
450 return *(this->self()); \
452 static_assert(true, \
453 "This assert is used to counter the false positive extra " \
454 "semi-colon warnings")
469 #define GKO_FACTORY_PARAMETER_SCALAR(_name, _default) \
470 GKO_FACTORY_PARAMETER(_name, _default)
485 #define GKO_FACTORY_PARAMETER_VECTOR(_name, ...) \
486 GKO_FACTORY_PARAMETER(_name, __VA_ARGS__)
487 #else // defined(__CUDACC__) || defined(__HIPCC__)
492 #define GKO_FACTORY_PARAMETER(_name, ...) \
493 _name{__VA_ARGS__}; \
495 template <typename... Args> \
496 auto with_##_name(Args&&... _value) \
497 ->std::decay_t<decltype(*(this->self()))>& \
499 GKO_NOT_IMPLEMENTED; \
500 return *(this->self()); \
502 static_assert(true, \
503 "This assert is used to counter the false positive extra " \
504 "semi-colon warnings")
506 #define GKO_FACTORY_PARAMETER_SCALAR(_name, _default) \
509 template <typename Arg> \
510 auto with_##_name(Arg&& _value)->std::decay_t<decltype(*(this->self()))>& \
512 using type = decltype(this->_name); \
513 this->_name = type{std::forward<Arg>(_value)}; \
514 return *(this->self()); \
516 static_assert(true, \
517 "This assert is used to counter the false positive extra " \
518 "semi-colon warnings")
520 #define GKO_FACTORY_PARAMETER_VECTOR(_name, ...) \
521 _name{__VA_ARGS__}; \
523 template <typename... Args> \
524 auto with_##_name(Args&&... _value) \
525 ->std::decay_t<decltype(*(this->self()))>& \
527 using type = decltype(this->_name); \
528 this->_name = type{std::forward<Args>(_value)...}; \
529 return *(this->self()); \
531 static_assert(true, \
532 "This assert is used to counter the false positive extra " \
533 "semi-colon warnings")
534 #endif // defined(__CUDACC__) || defined(__HIPCC__)
545 #define GKO_DEFERRED_FACTORY_PARAMETER(_name) \
549 using _name##_type = ::gko::detail::element_type_t<decltype(_name)>; \
552 auto with_##_name(::gko::deferred_factory_parameter<_name##_type> factory) \
553 ->std::decay_t<decltype(*(this->self()))>& \
555 this->_name##_generator_ = std::move(factory); \
556 this->deferred_factories[#_name] = [](const auto& exec, \
558 if (!params._name##_generator_.is_empty()) { \
559 params._name = params._name##_generator_.on(exec); \
562 return *(this->self()); \
566 ::gko::deferred_factory_parameter<_name##_type> _name##_generator_; \
569 static_assert(true, \
570 "This assert is used to counter the false positive extra " \
571 "semi-colon warnings")
583 #define GKO_DEFERRED_FACTORY_VECTOR_PARAMETER(_name) \
587 using _name##_type = ::gko::detail::value_element_type_t<decltype(_name)>; \
590 template <typename... Args, \
591 typename = std::enable_if_t<::std::conjunction< \
592 std::is_convertible<Args, ::gko::deferred_factory_parameter< \
593 _name##_type>>...>::value>> \
594 auto with_##_name(Args&&... factories) \
595 ->std::decay_t<decltype(*(this->self()))>& \
597 this->_name##_generator_ = { \
598 ::gko::deferred_factory_parameter<_name##_type>{ \
599 std::forward<Args>(factories)}...}; \
600 this->deferred_factories[#_name] = [](const auto& exec, \
602 if (!params._name##_generator_.empty()) { \
603 params._name.clear(); \
604 for (auto& generator : params._name##_generator_) { \
605 params._name.push_back(generator.on(exec)); \
609 return *(this->self()); \
611 template <typename FactoryType, \
612 typename = std::enable_if_t<std::is_convertible< \
614 ::gko::deferred_factory_parameter<_name##_type>>::value>> \
615 auto with_##_name(const std::vector<FactoryType>& factories) \
616 ->std::decay_t<decltype(*(this->self()))>& \
618 this->_name##_generator_.clear(); \
619 for (const auto& factory : factories) { \
620 this->_name##_generator_.push_back(factory); \
622 this->deferred_factories[#_name] = [](const auto& exec, \
624 if (!params._name##_generator_.empty()) { \
625 params._name.clear(); \
626 for (auto& generator : params._name##_generator_) { \
627 params._name.push_back(generator.on(exec)); \
631 return *(this->self()); \
635 std::vector<::gko::deferred_factory_parameter<_name##_type>> \
636 _name##_generator_; \
639 static_assert(true, \
640 "This assert is used to counter the false positive extra " \
641 "semi-colon warnings")
647 #endif // GKO_PUBLIC_CORE_BASE_ABSTRACT_FACTORY_HPP_