17 if (ad_ty)
return ad_ty;
31 auto& world = type->world();
37 auto T = callee->as<
App>()->arg();
38 auto [a, b] = arg->
projs<2>();
43 if (
auto sig = T->isa<
Sigma>()) {
44 auto p = sig->num_ops();
45 auto ops =
DefVec(p, [&](
size_t i) {
46 return world.app(world.app(world.annex<
add>(), sig->op(i)), {a->proj(i), b->proj(i)});
48 return world.tuple(ops);
49 }
else if (
auto arr = T->isa<
Arr>()) {
51 auto pack = world.mut_pack(T);
52 auto body_type = arr->body();
53 pack->set(world.app(world.app(world.annex<
add>(), body_type),
54 {world.extract(a, pack->var()), world.extract(b, pack->var())}));
61 }
else if (T->isa<
App>()) {
62 assert(0 &&
"not handled");
69 auto& world = type->world();
71 auto [count, T] = callee->as<
App>()->args<2>();
73 if (
auto lit = count->isa<
Lit>()) {
74 auto val = lit->get<
nat_t>();
75 auto args = arg->
projs(val);
76 auto sum = world.app(world.annex<
zero>(), T);
78 if (val >= 1)
sum = args[0];
79 for (
size_t i = 1; i < val; ++i)
80 sum = world.app(world.app(world.annex<
add>(), T), {sum, args[i]});
#define MIM_autodiff_NORMALIZER_IMPL
A (possibly paramterized) Array.
static auto isa(const Def *def)
auto projs(F f) const
Splits this Def via Def::projections into an Array (if A == std::dynamic_extent) or std::array (other...
static const Def * isa(const Def *def)
Checks if def is a Idx s and returns s or nullptr otherwise.
The automatic differentiation Plugin
const Def * normalize_Tangent(const Def *, const Def *, const Def *arg)
const Def * normalize_add(const Def *type, const Def *callee, const Def *arg)
Currently resolved the full addition.
const Def * autodiff_type_fun(const Def *)
const Def * normalize_AD(const Def *, const Def *, const Def *arg)
const Def * normalize_ad(const Def *, const Def *, const Def *)
Currently this normalizer does nothin.
const Def * tangent_type_fun(const Def *)
const Def * normalize_sum(const Def *type, const Def *callee, const Def *arg)
const Def * normalize_zero(const Def *, const Def *, const Def *)
Currently this normalizer does nothing.
Vector< const Def * > DefVec