MimIR
MimIR is my Intermediate Representation
Loading...
Searching...
No Matches
mim::plug::autodiff::phase::Eval Class Reference

This phase is the heart of AD. More...

#include <mim/plug/autodiff/phase/eval.h>

Inheritance diagram for mim::plug::autodiff::phase::Eval:
[legend]

Public Member Functions

 Eval (World &world, flags_t annex)
const Defrewrite_imm_App (const App *) final
 Detect autodiff calls.
const Defderive (const Def *)
 Acts on toplevel autodiff on closed terms:
const Defderive_ (const Def *)
 Additionally to the derivation, the pullback is registered and the maps are initialized.
const Defaugment (const Def *, Lam *, Lam *)
 Applies to (open) expressions in a functional context.
const Defaugment_ (const Def *, Lam *, Lam *)
 Rewrites the given definition in a lambda environment.
const Defaugment_var (const Var *, Lam *, Lam *)
 helper functions for augment
const Defaugment_lam (Lam *, Lam *, Lam *)
const Defaugment_extract (const Extract *, Lam *, Lam *)
const Defaugment_app (const App *, Lam *, Lam *)
const Defaugment_lit (const Lit *, Lam *, Lam *)
const Defaugment_tuple (const Tuple *, Lam *, Lam *)
const Defaugment_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)
Analysisanalysis ()
const Analysisanalysis () const
const Deflattice (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 Defabstracted (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).
Worldworld ()=delete
 Hides both and forbids direct access.
Worldold_world ()
 Get old Defs from here.
Worldnew_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< Phaserecreate ()
 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< Phasetake_resolved ()
 The Phase to use instead; nullptr means elide.
Worldworld ()
Driverdriver ()
Loglog () 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 ()
Worldworld ()
virtual void push ()
virtual void pop ()
virtual const Defmap (const Def *old_def, const Def *new_def)
const Defmap (const Def *old_def, Defs new_defs)
const Defmap (Defs old_defs, const Def *new_def)
const Defmap (Defs old_defs, Defs new_defs)
virtual const Deflookup (const Def *old_def)
 Lookup old_def by searching in reverse through the stack of maps.
virtual const Defrewrite (const Def *)
virtual const Defrewrite_imm (const Def *)
virtual const Defrewrite_mut (Def *)
virtual const Defrewrite_stub (Def *, Def *)
virtual DefVec rewrite (Defs)
virtual const Defrewrite_imm_Seq (const Seq *seq)
virtual const Defrewrite_mut_Seq (Seq *seq)

Additional Inherited Members

Static Public Member Functions inherited from mim::Phase
static std::unique_ptr< Phasecreate (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< Def2Defold2news_

Detailed Description

This phase is the heart of AD.

We replace an autodiff fun call with the differentiated function.

Definition at line 10 of file eval.h.

Constructor & Destructor Documentation

◆ Eval()

mim::plug::autodiff::phase::Eval::Eval ( World & world,
flags_t annex )
inline

Definition at line 12 of file eval.h.

References mim::Phase::annex(), and mim::RWPhase::RWPhase().

Member Function Documentation

◆ augment()

const Def * mim::plug::autodiff::phase::Eval::augment ( const Def * def,
Lam * f,
Lam * f_diff )

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_().

◆ augment_()

const Def * mim::plug::autodiff::phase::Eval::augment_ ( const Def * def,
Lam * f,
Lam * f_diff )

◆ augment_app()

◆ augment_extract()

const Def * mim::plug::autodiff::phase::Eval::augment_extract ( const Extract * ext,
Lam * f,
Lam * f_diff )

◆ augment_lam()

◆ augment_lit()

const Def * mim::plug::autodiff::phase::Eval::augment_lit ( const Lit * lit,
Lam * f,
Lam *  )

Definition at line 108 of file eval.cpp.

References mim::plug::autodiff::zero_pullback().

Referenced by augment_().

◆ augment_pack()

◆ augment_tuple()

◆ augment_var()

const Def * mim::plug::autodiff::phase::Eval::augment_var ( const Var * var,
Lam * ,
Lam *  )

helper functions for augment

Definition at line 114 of file eval.cpp.

Referenced by augment_().

◆ derive()

const Def * mim::plug::autodiff::phase::Eval::derive ( const Def * def)

Acts on toplevel autodiff on closed terms:

  • Replaces lambdas, operators with the appropriate derivatives.
  • Creates new lambda, calls associate variables, init maps, calls augment.

Definition at line 20 of file eval.cpp.

References derive_().

Referenced by augment_lam(), and rewrite_imm_App().

◆ derive_()

const Def * mim::plug::autodiff::phase::Eval::derive_ ( const Def * def)

◆ rewrite_imm_App()

const Def * mim::plug::autodiff::phase::Eval::rewrite_imm_App ( const App * app)
final

Detect autodiff calls.

Definition at line 26 of file eval.cpp.

References derive(), mim::RWPhase::is_bootstrapping(), mim::Axm::isa(), and mim::Rewriter::rewrite().


The documentation for this class was generated from the following files:
  • include/mim/plug/autodiff/phase/eval.h
  • src/mim/plug/autodiff/phase/eval.cpp