10using namespace std::literals;
31 auto& world = A->
world();
33 auto id_pb = world.mut_lam(arg_pb_ty)->set(
"id_pb");
34 auto id_pb_scalar = id_pb->var(0uz)->set(
"s");
43 auto& world = A->
world();
46 auto pb = world.mut_lam(pb_ty)->set(
"zero_pb");
47 pb->app(
true, pb->var(1), world.call<
zero>(A_tangent));
60 auto& world = E->
world();
63 auto pb_ty = world.cn({tang_ret, world.cn(tang_arg)});
70 auto& world = arg->
world();
73 if (!aug_arg || !aug_ret)
return nullptr;
78 auto deriv_ty = world.cn({aug_arg, world.cn({aug_ret, pb_ty})});
84 auto& world = pi->
world();
88 auto ret = pi->
codom();
91 if (!aug_arg)
return nullptr;
93 if (!aug_ret)
return nullptr;
94 return world.pi(aug_arg, aug_ret);
98 auto [arg, ret_pi] = pi->doms<2>();
99 auto ret = ret_pi->as<
Pi>()->dom();
106 auto& world = ty->
world();
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();
116 if (!body_ad)
return nullptr;
117 return world.arr(shape, body_ad);
119 if (
auto sig = ty->isa<
Sigma>()) {
121 auto ops =
DefVec(sig->ops(), [&](
const Def* op) { return autodiff_type_fun(op); });
122 return world.sigma(ops);
126 world.WLOG(
"no-diff type: {}", ty);
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);
142 auto zero = world.lit(T, 0)->set(
"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);
153 auto& world = T->
world();
154 return world.
app(world.app(world.annex<
sum>(), {world.lit_nat(defs.size()), T}), defs);
void reg_phases(Flags2Phases &phases)
void reg_phases(Flags2Phases &phases)
A (possibly paramterized) Array.
static auto isa(const Def *def)
World & world() const noexcept
static const Def * isa(const Def *def)
Checks if def is a Idx s and returns s or nullptr otherwise.
static void hook(Flags2Phases &phases)
A dependent function type.
static const Pi * isa_cn(const Def *d)
Is this a continuation - i.e. is the Pi::codom mim::Bottom?
const Def * codom() const
const Def * app(const Def *callee, const Def *arg)
The automatic differentiation Plugin
const Pi * autodiff_type_fun_pi(const Pi *)
const Def * op_sum(const Def *T, Defs)
const Def * autodiff_type_fun(const Def *)
const Def * zero_def(const Def *T)
const Def * tangent_type_fun(const Def *)
const Def * zero_pullback(const Def *E, const Def *A)
const Def * id_pullback(const Def *)
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...
Vector< const Def * > DefVec
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.
mim::Plugin mim_get_plugin()
absl::flat_hash_map< flags_t, NormalizeFn > Normalizers
#define MIM_REPL(__phases, __annex,...)
Basic info and registration function pointer to be returned from a specific plugin.