MimIR
MimIR is my Intermediate Representation
Loading...
Searching...
No Matches
lower_for.cpp
Go to the documentation of this file.
2
3#include <mim/driver.h>
4#include <mim/lam.h>
5#include <mim/tuple.h>
6
7#include <mim/plug/ll/ll.h>
8#include <mim/plug/mem/mem.h>
9
11
13
14namespace {
15
16const Def* merge_s(const Def* elem, const Def* sigma, const Def* mem) {
17 auto& w = elem->world();
18 if (mem) {
19 auto elems = sigma->projs();
20 return cat_sigma(elem, elems);
21 }
22 return w.sigma({elem, sigma});
23}
24
25const Def* merge_t(const Def* elem, const Def* tuple, const Def* mem) {
26 auto& w = elem->world();
27 if (mem) {
28 auto elems = tuple->projs();
29 return cat_tuple(elem, elems);
30 }
31 return w.tuple({elem, tuple});
32}
33
34} // namespace
35
36const Def* LowerFor::rewrite_imm_App(const App* app) {
37 if (is_bootstrapping()) return RWPhase::rewrite_imm_App(app);
38
39 if (auto for_ax = Axm::isa<affine::For>(app)) {
40 log().d("lower for: {}", for_ax);
41 auto [old_body, old_exit, args] = for_ax->uncurry_args<3>();
42 auto [new_begin, new_end, new_step, new_init] = args->projs<4>([this](const Def* def) { return rewrite(def); });
43
44 const Def* vec_axm = nullptr;
45 if (auto ll_vec = Axm::isa<ll::vec>(old_body)) {
46 old_body = ll_vec->arg();
47 if (old_world().driver().is_loaded("ll")) vec_axm = ll_vec->axm();
48 }
49
50 auto old_body_lam = old_body->isa_mut<Lam>();
51 auto old_exit_lam = old_exit->isa_mut<Lam>();
52 if (!old_body_lam) old_body_lam = Lam::eta_expand(old_body);
53 if (!old_exit_lam) old_exit_lam = Lam::eta_expand(old_exit);
54
55 auto new_mem = mem::mem_def(new_init);
56 auto new_head_lam = new_world().mut_con(merge_s(new_begin->type(), new_init->type(), new_mem))->set("head");
57 auto new_phis = new_head_lam->vars();
58 auto new_iter = new_phis.front();
59 auto new_acc = new_world().tuple(new_phis.view().subspan(1));
60 new_mem = mem::mem_var(new_head_lam);
61 auto new_bb_dom = new_mem ? new_mem->type() : new_world().sigma();
62
63 auto new_body = new_world().mut_con(new_bb_dom)->set("new_body");
64 auto new_exit = new_world().mut_con(new_bb_dom)->set("new_exit");
65 auto new_yield = new_world().mut_con(new_init->type())->set("new_yield");
66 auto new_cmp = new_world().call(core::icmp::ul, Defs{new_iter, new_end});
67 auto new_inc = new_world().call(core::wrap::add, core::Mode::nsuw, Defs{new_iter, new_step});
68
69 if (vec_axm) new_cmp = new_world().app(new_world().app(rewrite(vec_axm), new_cmp->type()), new_cmp);
70
71 new_head_lam->branch(false, new_cmp, new_body, new_exit, new_mem);
72 new_yield->app(false, new_head_lam, merge_t(new_inc, new_yield->var(), new_mem));
73
74 // `new_acc` references the head's phis, including the head's mem var.
75 // Each new bb receives its own mem, so re-thread the acc's mem through the bb's own mem var.
76 auto acc_for = [&](Lam* bb) -> const Def* {
77 if (!new_mem) return new_acc;
78 auto bb_mem = mem::mem_var(bb);
79 auto elems = DefVec();
80 for (auto phi : new_phis.view().subspan(1))
81 elems.emplace_back(Axm::isa<mem::M>(phi->type()) ? bb_mem : phi);
82 return new_world().tuple(elems);
83 };
84
85 push();
86 map(old_body_lam->var(), {new_iter, acc_for(new_body), new_yield});
87 auto new_body_filter = rewrite(old_body_lam->filter());
88 auto new_body_value = rewrite(old_body_lam->body());
89 new_body->set({new_body_filter, new_body_value});
90 pop();
91
92 push();
93 map(old_exit_lam->var(), acc_for(new_exit));
94 auto new_exit_filter = rewrite(old_exit_lam->filter());
95 auto new_exit_value = rewrite(old_exit_lam->body());
96 new_exit->set({new_exit_filter, new_exit_value});
97 pop();
98
99 return new_world().app(new_head_lam, merge_t(new_begin, new_init, new_mem));
100 }
101
102 return RWPhase::rewrite_imm_App(app);
103}
104
105} // namespace mim::plug::affine::phase
static auto isa(const Def *def)
Definition axm.h:112
Base class for all Defs.
Definition def.h:273
T * isa_mut() const
If this is mutable, it will cast constness away and perform a dynamic_cast to T.
Definition def.h:580
const Def * var(nat_t a, nat_t i) noexcept
Definition def.h:479
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
Definition def.h:1111
auto vars(F f) noexcept
Definition def.h:479
A function.
Definition lam.h:113
const Def * filter() const
Definition lam.h:125
Lam * set(Filter filter, const Def *body)
Definition lam.cpp:27
static Lam * eta_expand(Filter, const Def *f)
Definition lam.cpp:56
const Def * body() const
Definition lam.h:126
const fe::Log & log() const
Definition phase.h:79
Driver & driver()
Definition phase.h:78
const fe::Vector< std::string > & args()
Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>.
Definition phase.cpp:23
bool is_bootstrapping() const
Returns whether we are currently bootstrapping (rewriting annexes).
Definition phase.h:403
World & new_world()
Create new Defs into this.
Definition phase.h:452
World & old_world()
Get old Defs from here.
Definition phase.h:451
virtual void push()
Definition rewrite.h:40
virtual const Def * map(const Def *old_def, const Def *new_def)
Definition rewrite.h:47
virtual void pop()
Definition rewrite.h:41
virtual const Def * rewrite(const Def *)
Definition rewrite.cpp:55
const Def * sigma(Defs ops)
Definition world.cpp:316
const Def * app(const Def *callee, const Def *arg)
Definition world.cpp:237
const Def * tuple(Defs ops)
Definition world.cpp:326
const Def * call(const Def *callee, T &&arg, Args &&... args)
Definition world.h:662
Lam * mut_con(const Def *dom)
Definition world.h:414
const Def * rewrite_imm_App(const App *) final
Definition lower_for.cpp:36
const Def * mem_var(Lam *lam)
Returns the memory argument of a function if it has one.
Definition mem.h:55
const Def * mem_def(const Def *def)
Returns the (first) element of type mem.M a from the given tuple.
Definition mem.h:42
const Def * cat_tuple(nat_t n, nat_t m, const Def *a, const Def *b)
Definition tuple.cpp:92
fe::View< const Def * > Defs
Definition def.h:91
const Def * cat_sigma(nat_t n, nat_t m, const Def *a, const Def *b)
Definition tuple.cpp:93
fe::Vector< const Def * > DefVec
Definition def.h:93