MimIR
MimIR is my Intermediate Representation
Loading...
Searching...
No Matches
autodiff.cpp
Go to the documentation of this file.
2
3#include <mim/config.h>
4#include <mim/phase.h>
5
6#include <mim/plug/mem/mem.h>
7
9
10using namespace std::literals;
11using namespace mim;
12using namespace mim::plug;
13
14void reg_phases(Flags2Phases& phases) {
16
18 if (auto zero = Axm::isa<autodiff::zero>(def); zero) {
19 if (auto z = autodiff::zero_def(zero->arg())) return z;
20 }
21 return {};
22 });
23}
24
26 return {"autodiff", MIM_VERSION, [](Normalizers& n) { autodiff::register_normalizers(n); }, reg_phases};
27}
28
29namespace mim::plug::autodiff {
30const Def* id_pullback(const Def* A) {
31 auto& world = A->world();
32 auto arg_pb_ty = pullback_type(A, A);
33 auto id_pb = world.mut_lam(arg_pb_ty)->set("id_pb");
34 auto id_pb_scalar = id_pb->var(0uz)->set("s");
35 id_pb->app(true,
36 id_pb->var(1), // can not use ret_var as the result might be higher order
37 id_pb_scalar);
38
39 return id_pb;
40}
41
42const Def* zero_pullback(const Def* E, const Def* A) {
43 auto& world = A->world();
44 auto A_tangent = tangent_type_fun(A);
45 auto pb_ty = pullback_type(E, A);
46 auto pb = world.mut_lam(pb_ty)->set("zero_pb");
47 pb->app(true, pb->var(1), world.call<zero>(A_tangent));
48 return pb;
49}
50
51// `P` => `P*`
52// TODO: nothing? function => R? Mem => R?
53// TODO: rename to op_tangent_type
54const Def* tangent_type_fun(const Def* ty) { return ty; }
55
56/// computes pb type `E* -> A*`
57/// `E` - type of the expression (return type for a function)
58/// `A` - type of the argument (point of orientation resp. derivative - argument type for partial pullbacks)
59const Pi* pullback_type(const Def* E, const Def* A) {
60 auto& world = E->world();
61 auto tang_arg = tangent_type_fun(A);
62 auto tang_ret = tangent_type_fun(E);
63 auto pb_ty = world.cn({tang_ret, world.cn(tang_arg)});
64 return pb_ty;
65}
66
67namespace {
68// `A,R` => `(A->R)' = A' -> R' * (R* -> A*)`
69const Pi* autodiff_type_fun(const Def* arg, const Def* ret) {
70 auto& world = arg->world();
73 if (!aug_arg || !aug_ret) return nullptr;
74 // `Q* -> P*`
75 auto pb_ty = pullback_type(ret, arg);
76 // `P' -> Q' * (Q* -> P*)`
77
78 auto deriv_ty = world.cn({aug_arg, world.cn({aug_ret, pb_ty})});
79 return deriv_ty;
80}
81} // namespace
82
83const Pi* autodiff_type_fun_pi(const Pi* pi) {
84 auto& world = pi->world();
85 if (!Pi::isa_cn(pi)) {
86 // TODO: dependency
87 auto arg = pi->dom();
88 auto ret = pi->codom();
89 if (ret->isa<Pi>()) {
90 auto aug_arg = autodiff_type_fun(arg);
91 if (!aug_arg) return nullptr;
92 auto aug_ret = autodiff_type_fun(pi->codom());
93 if (!aug_ret) return nullptr;
94 return world.pi(aug_arg, aug_ret);
95 }
96 return autodiff_type_fun(arg, ret);
97 }
98 auto [arg, ret_pi] = pi->doms<2>();
99 auto ret = ret_pi->as<Pi>()->dom();
100 return autodiff_type_fun(arg, ret);
101}
102
103// In general transforms `A` => `A'`.
104// Especially `P->Q` => `P'->Q' * (Q* -> P*)`.
105const Def* autodiff_type_fun(const Def* ty) {
106 auto& world = ty->world();
107 // TODO: handle DS (operators)
108 if (auto pi = ty->isa<Pi>()) return autodiff_type_fun_pi(pi);
109 // Also handles autodiff call from axm declaration => abstract => leave it.
110 if (Idx::isa(ty)) return ty;
111 if (ty == world.type_nat()) return ty;
112 if (auto arr = ty->isa<Arr>()) {
113 auto shape = arr->arity();
114 auto body = arr->body();
115 auto body_ad = autodiff_type_fun(body);
116 if (!body_ad) return nullptr;
117 return world.arr(shape, body_ad);
118 }
119 if (auto sig = ty->isa<Sigma>()) {
120 // TODO: mut sigma
121 auto ops = DefVec(sig->ops(), [&](const Def* op) { return autodiff_type_fun(op); });
122 return world.sigma(ops);
123 }
124 // mem
125 if (Axm::isa<mem::M>(ty)) return ty;
126 world.WLOG("no-diff type: {}", ty);
127 return nullptr;
128}
129
130const Def* zero_def(const Def* T) {
131 // TODO: we want: zero mem -> zero mem or bot
132 // zero [A,B,C] -> [zero A, zero B, zero C]
133 auto& world = T->world();
134 if (auto arr = T->isa<Arr>()) {
135 auto arity = arr->arity();
136 auto body = arr->body();
137 auto inner_zero = world.app(world.annex<zero>(), body);
138 auto zero_arr = world.pack(arity, inner_zero);
139 return zero_arr;
140 } else if (Idx::isa(T)) {
141 // TODO: real
142 auto zero = world.lit(T, 0)->set("zero");
143 return zero;
144 } else if (auto sig = T->isa<Sigma>()) {
145 auto ops = DefVec(sig->ops(), [&](const Def* op) { return world.app(world.annex<zero>(), op); });
146 return world.tuple(ops);
147 }
148 return nullptr;
149}
150
151const Def* op_sum(const Def* T, Defs defs) {
152 // TODO: assert all are of type T
153 auto& world = T->world();
154 return world.app(world.app(world.annex<sum>(), {world.lit_nat(defs.size()), T}), defs);
155}
156
157} // namespace mim::plug::autodiff
void reg_phases(Flags2Phases &phases)
Definition affine.cpp:12
void reg_phases(Flags2Phases &phases)
Definition autodiff.cpp:14
A (possibly paramterized) Array.
Definition tuple.h:121
static auto isa(const Def *def)
Definition axm.h:107
Base class for all Defs.
Definition def.h:261
World & world() const noexcept
Definition def.cpp:483
static const Def * isa(const Def *def)
Checks if def is a Idx s and returns s or nullptr otherwise.
Definition def.cpp:658
static void hook(Flags2Phases &phases)
Definition phase.h:70
A dependent function type.
Definition lam.h:14
static const Pi * isa_cn(const Def *d)
Is this a continuation - i.e. is the Pi::codom mim::Bottom?
Definition lam.h:47
const Def * dom() const
Definition lam.h:35
const Def * codom() const
Definition lam.h:36
A dependent tuple type.
Definition tuple.h:22
const Def * app(const Def *callee, const Def *arg)
Definition world.cpp:205
#define MIM_EXPORT
Definition config.h:19
The automatic differentiation Plugin
Definition autodiff.h:7
const Pi * autodiff_type_fun_pi(const Pi *)
Definition autodiff.cpp:83
const Def * op_sum(const Def *T, Defs)
Definition autodiff.cpp:151
const Def * autodiff_type_fun(const Def *)
Definition autodiff.cpp:105
const Def * zero_def(const Def *T)
Definition autodiff.cpp:130
const Def * tangent_type_fun(const Def *)
Definition autodiff.cpp:54
const Def * zero_pullback(const Def *E, const Def *A)
Definition autodiff.cpp:42
const Def * id_pullback(const Def *)
Definition autodiff.cpp:30
void register_normalizers(Normalizers &normalizers)
const Pi * pullback_type(const Def *E, const Def *A)
computes pb type E* -> A* E - type of the expression (return type for a function) A - type of the arg...
Definition autodiff.cpp:59
Definition ast.h:14
View< const Def * > Defs
Definition def.h:78
Vector< const Def * > DefVec
Definition def.h:79
absl::flat_hash_map< flags_t, std::function< std::unique_ptr< Phase >(World &)> > Flags2Phases
Maps an axiom of a Phase to a function that creates one.
Definition plugin.h:25
mim::Plugin mim_get_plugin()
absl::flat_hash_map< flags_t, NormalizeFn > Normalizers
Definition plugin.h:22
#define MIM_REPL(__phases, __annex,...)
Definition phase.h:406
#define MIM_VERSION
Definition plugin.h:54
Basic info and registration function pointer to be returned from a specific plugin.
Definition plugin.h:59