5 #ifndef GKO_PUBLIC_CORE_LOG_LOGGER_HPP_
6 #define GKO_PUBLIC_CORE_LOG_LOGGER_HPP_
12 #include <type_traits>
15 #include <ginkgo/core/base/types.hpp>
16 #include <ginkgo/core/base/utils_helper.hpp>
22 template <
typename ValueType>
27 class PolymorphicObject;
29 class stopping_status;
31 class EnableCloneable;
38 class BatchLinOpFactory;
40 template <
typename ValueType>
112 #define GKO_LOGGER_REGISTER_EVENT(_id, _event_name, ...) \
114 virtual void on_##_event_name(__VA_ARGS__) const {} \
117 template <size_type Event, typename... Params> \
118 std::enable_if_t<Event == _id && (_id < event_count_max)> on( \
119 Params&&... params) const \
121 if (enabled_events_ & (mask_type{1} << _id)) { \
122 this->on_##_event_name(std::forward<Params>(params)...); \
125 static constexpr size_type _event_name{_id}; \
126 static constexpr mask_type _event_name##_mask{mask_type{1} << _id};
134 GKO_LOGGER_REGISTER_EVENT(0, allocation_started,
const Executor* exec,
144 GKO_LOGGER_REGISTER_EVENT(1, allocation_completed,
const Executor* exec,
154 GKO_LOGGER_REGISTER_EVENT(2, free_started,
const Executor* exec,
163 GKO_LOGGER_REGISTER_EVENT(3, free_completed,
const Executor* exec,
175 GKO_LOGGER_REGISTER_EVENT(4, copy_started,
const Executor* exec_from,
188 GKO_LOGGER_REGISTER_EVENT(5, copy_completed,
const Executor* exec_from,
198 GKO_LOGGER_REGISTER_EVENT(6, operation_launched,
const Executor* exec,
212 GKO_LOGGER_REGISTER_EVENT(7, operation_completed,
const Executor* exec,
221 GKO_LOGGER_REGISTER_EVENT(8, polymorphic_object_create_started,
231 GKO_LOGGER_REGISTER_EVENT(9, polymorphic_object_create_completed,
243 GKO_LOGGER_REGISTER_EVENT(10, polymorphic_object_copy_started,
255 GKO_LOGGER_REGISTER_EVENT(11, polymorphic_object_copy_completed,
266 GKO_LOGGER_REGISTER_EVENT(12, polymorphic_object_deleted,
276 GKO_LOGGER_REGISTER_EVENT(13, linop_apply_started,
const LinOp* A,
286 GKO_LOGGER_REGISTER_EVENT(14, linop_apply_completed,
const LinOp* A,
298 GKO_LOGGER_REGISTER_EVENT(15, linop_advanced_apply_started,
const LinOp* A,
311 GKO_LOGGER_REGISTER_EVENT(16, linop_advanced_apply_completed,
322 GKO_LOGGER_REGISTER_EVENT(17, linop_factory_generate_started,
333 GKO_LOGGER_REGISTER_EVENT(18, linop_factory_generate_completed,
348 GKO_LOGGER_REGISTER_EVENT(19, criterion_check_started,
352 const uint8& stopping_id,
353 const bool& set_finalized)
375 GKO_LOGGER_REGISTER_EVENT(
378 const uint8& stopping_id,
const bool& set_finalized,
380 const bool& all_converged)
399 virtual void on_criterion_check_completed(
402 const uint8& stopping_id,
const bool& set_finalized,
404 const bool& all_converged)
const
406 this->on_criterion_check_completed(
criterion, it, r, tau, x,
407 stopping_id, set_finalized, status,
408 one_changed, all_converged);
412 static constexpr
size_type iteration_complete{21};
413 static constexpr mask_type iteration_complete_mask{mask_type{1} << 21};
415 template <
size_type Event,
typename... Params>
417 Params&&... params)
const
419 if (enabled_events_ & (mask_type{1} << 21)) {
420 this->on_iteration_complete(std::forward<Params>(params)...);
439 "Please use the version with the additional stopping "
441 virtual
void on_iteration_complete(const LinOp*
solver, const
size_type& it,
442 const LinOp* r, const LinOp* x =
nullptr,
443 const LinOp* tau =
nullptr)
const
461 "Please use the version with the additional stopping "
463 virtual
void on_iteration_complete(const LinOp*
solver, const
size_type& it,
464 const LinOp* r, const LinOp* x,
466 const LinOp* implicit_tau_sq)
const
468 GKO_BEGIN_DISABLE_DEPRECATION_WARNINGS
469 this->on_iteration_complete(
solver, it, r, x, tau);
470 GKO_END_DISABLE_DEPRECATION_WARNINGS
488 virtual void on_iteration_complete(
const LinOp*
solver,
const LinOp* b,
490 const LinOp* r,
const LinOp* tau,
491 const LinOp* implicit_tau_sq,
492 const array<stopping_status>* status,
495 GKO_BEGIN_DISABLE_DEPRECATION_WARNINGS
496 this->on_iteration_complete(
solver, it, r, x, tau, implicit_tau_sq);
497 GKO_END_DISABLE_DEPRECATION_WARNINGS
508 GKO_LOGGER_REGISTER_EVENT(22, polymorphic_object_move_started,
509 const Executor* exec,
510 const PolymorphicObject* input,
511 const PolymorphicObject* output)
520 GKO_LOGGER_REGISTER_EVENT(23, polymorphic_object_move_completed,
521 const Executor* exec,
522 const PolymorphicObject* input,
523 const PolymorphicObject* output)
532 GKO_LOGGER_REGISTER_EVENT(24, batch_linop_factory_generate_started,
533 const batch::BatchLinOpFactory*
factory,
534 const batch::BatchLinOp* input)
544 GKO_LOGGER_REGISTER_EVENT(25, batch_linop_factory_generate_completed,
545 const batch::BatchLinOpFactory*
factory,
546 const batch::BatchLinOp* input,
547 const batch::BatchLinOp* output)
550 static constexpr
size_type batch_solver_completed{26};
551 static constexpr mask_type batch_solver_completed_mask{mask_type{1} << 26};
553 template <
size_type Event,
typename... Params>
555 Params&&... params)
const
557 if (enabled_events_ & batch_solver_completed_mask) {
558 this->on_batch_solver_completed(std::forward<Params>(params)...);
570 virtual void on_batch_solver_completed(
571 const array<int>& iters,
const array<double>& residual_norms)
const
581 virtual void on_batch_solver_completed(
582 const array<int>& iters,
const array<float>& residual_norms)
const
586 #if GINKGO_ENABLE_HALF
596 virtual void on_batch_solver_completed(
597 const array<int>& iters,
598 const array<gko::float16>& residual_norms)
const
605 #if GINKGO_ENABLE_BFLOAT16
615 virtual void on_batch_solver_completed(
616 const array<int>& iters,
617 const array<gko::bfloat16>& residual_norms)
const
625 #undef GKO_LOGGER_REGISTER_EVENT
631 allocation_started_mask | allocation_completed_mask |
632 free_started_mask | free_completed_mask | copy_started_mask |
639 operation_launched_mask | operation_completed_mask;
645 polymorphic_object_create_started_mask |
646 polymorphic_object_create_completed_mask |
647 polymorphic_object_copy_started_mask |
648 polymorphic_object_copy_completed_mask |
649 polymorphic_object_move_started_mask |
650 polymorphic_object_move_completed_mask |
651 polymorphic_object_deleted_mask;
657 linop_apply_started_mask | linop_apply_completed_mask |
658 linop_advanced_apply_started_mask | linop_advanced_apply_completed_mask;
664 linop_factory_generate_started_mask |
665 linop_factory_generate_completed_mask;
671 batch_linop_factory_generate_started_mask |
672 batch_linop_factory_generate_completed_mask;
678 criterion_check_started_mask | criterion_check_completed_mask;
686 virtual ~
Logger() =
default;
703 GKO_DEPRECATED(
"use single-parameter constructor")
724 : enabled_events_{enabled_events}
728 mask_type enabled_events_;
746 virtual void add_logger(std::shared_ptr<const Logger> logger) = 0;
769 virtual const std::vector<std::shared_ptr<const Logger>>&
get_loggers()
789 template <
typename ConcreteLoggable,
typename PolymorphicBase = Loggable>
791 template <
typename T>
795 void add_logger(std::shared_ptr<const Logger> logger)
override
797 loggers_.push_back(logger);
800 void remove_logger(
const Logger* logger)
override
803 find_if(begin(loggers_), end(loggers_),
804 [&logger](
const auto& l) {
return l.get() == logger; });
805 if (idx != end(loggers_)) {
815 remove_logger(logger.
get());
818 const std::vector<std::shared_ptr<const Logger>>& get_loggers()
824 void clear_loggers()
override { loggers_.clear(); }
834 template <
size_type Event,
typename ConcreteLoggableT,
typename =
void>
835 struct propagate_log_helper {
836 template <
typename... Args>
837 static void propagate_log(
const ConcreteLoggableT*, Args&&...)
841 template <
size_type Event,
typename ConcreteLoggableT>
842 struct propagate_log_helper<
843 Event, ConcreteLoggableT,
845 decltype(std::declval<ConcreteLoggableT>().get_executor())>> {
846 template <
typename... Args>
847 static void propagate_log(
const ConcreteLoggableT* loggable,
850 const auto exec = loggable->get_executor();
851 if (exec->should_propagate_log()) {
852 for (
auto& logger : exec->get_loggers()) {
853 if (logger->needs_propagation()) {
854 logger->template on<Event>(std::forward<Args>(args)...);
862 template <
size_type Event,
typename... Params>
863 void log(Params&&... params)
const
865 propagate_log_helper<Event, ConcreteLoggable>::propagate_log(
866 static_cast<const ConcreteLoggable*>(
this),
867 std::forward<Params>(params)...);
868 for (
auto& logger : loggers_) {
869 logger->template on<Event>(std::forward<Params>(params)...);
873 std::vector<std::shared_ptr<const Logger>> loggers_;
881 #endif // GKO_PUBLIC_CORE_LOG_LOGGER_HPP_