This phase is the heart of AD. More...
#include <mim/plug/autodiff/phase/eval.h>
Public Member Functions | |
| Eval (World &world, flags_t annex) | |
| const Def * | rewrite_imm_App (const App *) final |
| Detect autodiff calls. | |
| const Def * | derive (const Def *) |
| Acts on toplevel autodiff on closed terms: | |
| const Def * | derive_ (const Def *) |
| Additionally to the derivation, the pullback is registered and the maps are initialized. | |
| const Def * | augment (const Def *, Lam *, Lam *) |
| Applies to (open) expressions in a functional context. | |
| const Def * | augment_ (const Def *, Lam *, Lam *) |
| Rewrites the given definition in a lambda environment. | |
| const Def * | augment_var (const Var *, Lam *, Lam *) |
| helper functions for augment | |
| const Def * | augment_lam (Lam *, Lam *, Lam *) |
| const Def * | augment_extract (const Extract *, Lam *, Lam *) |
| const Def * | augment_app (const App *, Lam *, Lam *) |
| const Def * | augment_lit (const Lit *, Lam *, Lam *) |
| const Def * | augment_tuple (const Tuple *, Lam *, Lam *) |
| const Def * | augment_pack (const Pack *pack, Lam *f, Lam *f_diff) |
| Public Member Functions inherited from mim::RWPhase | |
| RWPhase (World &world, std::string name, Analysis *analysis=nullptr) | |
| RWPhase (World &world, flags_t annex, Analysis *analysis=nullptr) | |
| Analysis * | analysis () |
| const Analysis * | analysis () const |
| const Def * | lattice (const Def *old_def) const |
| Returns the abstract value computed by the associated Analysis for the given old-world Def, or nullptr if no value is available. | |
| const Def * | abstracted (const Def *old_def) const |
Returns lattice(old_def) if it differs from old_def (i.e. we learned something), otherwise nullptr. | |
| virtual bool | analyze () |
| Runs the optional pre-analysis on RWPhase::old_world(), typically to a fixed point, before rewriting begins. | |
| virtual void | rewrite_annex (flags_t, Sym, const Def *) |
| virtual void | rewrite_external (Def *) |
| bool | is_bootstrapping () const |
| Returns whether we are currently bootstrapping (rewriting annexes). | |
| World & | world ()=delete |
| Hides both and forbids direct access. | |
| World & | old_world () |
| Get old Defs from here. | |
| World & | new_world () |
| Create new Defs into this. | |
| Public Member Functions inherited from mim::Phase | |
| Phase (World &world, std::string name) | |
| Phase (World &world, flags_t annex) | |
| virtual | ~Phase ()=default |
| virtual std::unique_ptr< Phase > | recreate () |
| Creates a new instance; needed by a fixed-point PhaseMan. | |
| virtual void | apply (const App *) |
| Invoked if your Phase has additional args. | |
| virtual void | apply (Phase &) |
| Dito, but invoked by Phase::recreate. | |
| virtual bool | redirects () const |
| If true, Phase::create uses take_resolved(). | |
| virtual std::unique_ptr< Phase > | take_resolved () |
| The Phase to use instead; nullptr means elide. | |
| World & | world () |
| Driver & | driver () |
| Log & | log () const |
| std::string_view | name () const |
| flags_t | annex () const |
| const Vector< std::string > & | args () |
| Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>. | |
| bool | todo () const |
| void | invalidate (bool todo=true) |
| Signals that another round of fixed-point iteration is required, either as part of. | |
| virtual void | run () |
| Entry point and generates some debug output; invokes Phase::start. | |
| void | profile_count (std::string_view key, uint64_t n=1) |
Adds n to the custom Profiler counter key of the current run; no-op unless profiling is enabled. | |
| Public Member Functions inherited from mim::Rewriter | |
| template<class D = Def> | |
| D * | curr_mut () const |
| Rewriter (std::unique_ptr< World > &&ptr) | |
| Rewriter (World &world) | |
| virtual | ~Rewriter () |
| void | reset (std::unique_ptr< World > &&ptr) |
| void | reset () |
| World & | world () |
| virtual void | push () |
| virtual void | pop () |
| virtual const Def * | map (const Def *old_def, const Def *new_def) |
| const Def * | map (const Def *old_def, Defs new_defs) |
| const Def * | map (Defs old_defs, const Def *new_def) |
| const Def * | map (Defs old_defs, Defs new_defs) |
| virtual const Def * | lookup (const Def *old_def) |
| Lookup old_def by searching in reverse through the stack of maps. | |
| virtual const Def * | rewrite (const Def *) |
| virtual const Def * | rewrite_imm (const Def *) |
| virtual const Def * | rewrite_mut (Def *) |
| virtual const Def * | rewrite_stub (Def *, Def *) |
| virtual DefVec | rewrite (Defs) |
| virtual const Def * | rewrite_imm_Seq (const Seq *seq) |
| virtual const Def * | rewrite_mut_Seq (Seq *seq) |
Additional Inherited Members | |
| Static Public Member Functions inherited from mim::Phase | |
| static std::unique_ptr< Phase > | create (const Flags2Phases &phases, const Def *def) |
| template<class A, class P> | |
| static void | hook (Flags2Phases &phases) |
| template<class P, class... Args> | |
| static void | run (Args &&... args) |
| Runs a single Phase. | |
| Protected Member Functions inherited from mim::RWPhase | |
| void | start () override |
| Actual entry. | |
| Protected Member Functions inherited from mim::Rewriter | |
| auto | enter (Def *new_mut) |
Updates curr_mut() to new_mut and restores it at the end of the scope. | |
| Protected Attributes inherited from mim::Phase | |
| std::string | name_ |
| Protected Attributes inherited from mim::Rewriter | |
| std::deque< Def2Def > | old2news_ |
Definition at line 12 of file eval.h.
References mim::Phase::annex(), and mim::RWPhase::RWPhase().
Applies to (open) expressions in a functional context.
Returns the rewritten expressions and augments the partial and modular pullbacks. The rewrite is identity on the term up to renaming of variables. Otherwise, only pullbacks are added. To do so, some calls (e.g. axms) are replaced by their derivatives. This transformation can be seen as an augmentation with a dual computation that generates the derivatives.
Definition at line 14 of file eval.cpp.
References augment_().
Referenced by augment_app(), augment_extract(), augment_lam(), augment_pack(), augment_tuple(), and derive_().
Rewrites the given definition in a lambda environment.
Definition at line 349 of file eval.cpp.
References augment_app(), augment_extract(), augment_lam(), augment_lit(), augment_pack(), augment_tuple(), augment_var(), mim::plug::autodiff::autodiff_type_fun(), ELOG, mim::World::externals(), mim::find_and_replace(), mim::Def::isa_mut(), mim::Def::node_name(), mim::RWPhase::old_world(), mim::Rewriter::rewrite(), mim::World::sym(), and mim::Def::type().
Referenced by augment(), and augment_pack().
| const Def * mim::plug::autodiff::phase::Eval::augment_app | ( | const App * | app, |
| Lam * | f, | ||
| Lam * | f_diff ) |
Definition at line 269 of file eval.cpp.
References mim::World::app(), mim::App::arg(), augment(), mim::App::callee(), mim::compose_cn(), mim::Pi::isa_basicblock(), mim::Pi::isa_cn(), mim::World::mut_lam(), mim::plug::cps::op_cps2ds_dep(), mim::Def::projs(), mim::Lam::set(), mim::Def::type(), mim::Def::var(), and mim::RWPhase::world().
Referenced by augment_().
| const Def * mim::plug::autodiff::phase::Eval::augment_extract | ( | const Extract * | ext, |
| Lam * | f, | ||
| Lam * | f_diff ) |
Definition at line 171 of file eval.cpp.
References augment(), mim::World::extract(), mim::Extract::index(), mim::World::insert(), mim::World::mut_lam(), mim::plug::autodiff::pullback_type(), mim::Def::set(), mim::Lam::set(), mim::Extract::tuple(), mim::Def::type(), mim::Def::var(), and mim::RWPhase::world().
Referenced by augment_().
Definition at line 121 of file eval.cpp.
References augment(), mim::plug::autodiff::autodiff_type_fun(), mim::Lam::body(), derive(), mim::Pi::dom(), mim::Lam::filter(), mim::Lam::isa_basicblock(), mim::World::mut_con(), mim::plug::autodiff::pullback_type(), mim::Def::sym(), mim::Lam::type(), mim::Def::var(), mim::RWPhase::world(), and mim::plug::autodiff::zero_pullback().
Referenced by augment_().
Definition at line 108 of file eval.cpp.
References mim::plug::autodiff::zero_pullback().
Referenced by augment_().
| const Def * mim::plug::autodiff::phase::Eval::augment_pack | ( | const Pack * | pack, |
| Lam * | f, | ||
| Lam * | f_diff ) |
Definition at line 234 of file eval.cpp.
References mim::Phase::annex(), mim::World::app(), mim::Pack::arity(), augment(), augment_(), mim::Seq::body(), mim::World::mut_lam(), mim::World::mut_pack(), mim::plug::cps::op_cps2ds_dep(), mim::World::pack(), mim::plug::autodiff::pullback_type(), mim::Lam::set(), mim::Pack::set(), mim::plug::autodiff::tangent_type_fun(), mim::Def::type(), and mim::RWPhase::world().
Referenced by augment_().
| const Def * mim::plug::autodiff::phase::Eval::augment_tuple | ( | const Tuple * | tup, |
| Lam * | f, | ||
| Lam * | f_diff ) |
Definition at line 204 of file eval.cpp.
References augment(), mim::World::mut_lam(), mim::plug::autodiff::op_sum(), mim::Def::projs(), mim::plug::autodiff::pullback_type(), mim::Def::set(), mim::Lam::set(), mim::plug::autodiff::tangent_type_fun(), mim::World::tuple(), mim::Def::type(), mim::Def::var(), and mim::RWPhase::world().
Referenced by augment_().
Acts on toplevel autodiff on closed terms:
Definition at line 20 of file eval.cpp.
References derive_().
Referenced by augment_lam(), and rewrite_imm_App().
Additionally to the derivation, the pullback is registered and the maps are initialized.
Definition at line 52 of file eval.cpp.
References mim::Def::as_mut(), augment(), mim::plug::autodiff::autodiff_type_fun_pi(), mim::plug::autodiff::id_pullback(), mim::World::mut_lam(), mim::Def::set(), mim::Lam::set(), mim::World::tuple(), mim::Lam::type(), mim::Def::var(), mim::RWPhase::world(), and mim::plug::autodiff::zero_pullback().
Referenced by derive().
Detect autodiff calls.
Definition at line 26 of file eval.cpp.
References derive(), mim::RWPhase::is_bootstrapping(), mim::Axm::isa(), and mim::Rewriter::rewrite().