MimIR
0.3-dev
MimIR is my Intermediate Representation
Loading...
Searching...
No Matches
lower.cpp
Go to the documentation of this file.
1
#include "
mim/plug/tensor/phase/lower.h
"
2
3
#include "
mim/def.h
"
4
#include "
mim/lam.h
"
5
6
#include "
mim/plug/tensor/tensor.h
"
7
8
namespace
mim::plug::tensor::phase
{
9
10
const
Def* Lower::lower_via_impl(
const
App
* app,
const
Def* impl_annex) {
11
auto
&
w
=
new_world
();
12
13
// Walk the curry chain (innermost App outermost in syntax) to collect the args
14
// in the order they were applied.
15
DefVec
args
;
16
const
Def*
head
= app;
17
while
(
auto
h =
head
->isa<
App
>()) {
18
args
.push_back(
rewrite
(
h
->arg()));
19
head
=
h
->callee();
20
}
21
std::reverse(
args
.begin(),
args
.end());
22
23
auto
impl = impl_annex;
24
for
(
auto
a :
args
)
25
impl =
w
.app(impl, a);
26
27
// The `_impl` is a `lam`, so applying it triggers beta-reduction. Each `_impl`
28
// body references the `_impl` variants of its dependencies directly, so the
29
// chain bottoms out at the low-level axioms (`map_reduce`, …) in one go.
30
return
impl;
31
}
32
33
const
Def
*
Lower::rewrite_imm_App
(
const
App
* app) {
34
auto
& w =
new_world
();
35
36
if
(
Axm::isa<tensor::broadcast_in_dim>
(app))
37
return
lower_via_impl(app, w.annex<
tensor::broadcast_in_dim_impl
>());
38
else
if
(
Axm::isa<tensor::product_2d>
(app))
39
return
lower_via_impl(app, w.annex<
tensor::product_2d_impl
>());
40
else
if
(
Axm::isa<tensor::bmm>
(app))
41
return
lower_via_impl(app, w.annex<
tensor::bmm_impl
>());
42
else
if
(
Axm::isa<tensor::dot_product>
(app))
43
return
lower_via_impl(app, w.annex<
tensor::dot_product_impl
>());
44
else
if
(
Axm::isa<tensor::transpose>
(app))
45
return
lower_via_impl(app, w.annex<
tensor::transpose_impl
>());
46
else
if
(
Axm::isa<tensor::transpose_2d>
(app))
47
return
lower_via_impl(app, w.annex<
tensor::transpose_2d_impl
>());
48
else
if
(
Axm::isa<tensor::map>
(app))
49
return
lower_via_impl(app, w.annex<
tensor::map_impl
>());
50
else
if
(
Axm::isa<tensor::unary>
(app))
51
return
lower_via_impl(app, w.annex<
tensor::unary_impl
>());
52
else
if
(
Axm::isa<tensor::binary>
(app))
53
return
lower_via_impl(app, w.annex<
tensor::binary_impl
>());
54
else
if
(
Axm::isa<tensor::select>
(app))
55
return
lower_via_impl(app, w.annex<
tensor::select_impl
>());
56
else
if
(
Axm::isa<tensor::repeat>
(app))
57
return
lower_via_impl(app, w.annex<
tensor::repeat_impl
>());
58
else
if
(
Axm::isa<tensor::reshape>
(app))
59
return
lower_via_impl(app, w.annex<
tensor::reshape_impl
>());
60
else
if
(
Axm::isa<tensor::slice>
(app))
61
return
lower_via_impl(app, w.annex<
tensor::slice_impl
>());
62
else
if
(
Axm::isa<tensor::flip>
(app))
63
return
lower_via_impl(app, w.annex<
tensor::flip_impl
>());
64
else
if
(
Axm::isa<tensor::conv>
(app))
65
return
lower_via_impl(app, w.annex<
tensor::conv_impl
>());
66
else
if
(
Axm::isa<tensor::pool>
(app))
67
return
lower_via_impl(app, w.annex<
tensor::pool_impl
>());
68
return
RWPhase::rewrite_imm_App(app);
69
}
70
71
}
// namespace mim::plug::tensor::phase
mim::App
Definition
lam.h:224
mim::Axm::isa
static auto isa(const Def *def)
Definition
axm.h:107
mim::Def
Base class for all Defs.
Definition
def.h:261
mim::Phase::args
const Vector< std::string > & args()
Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>.
Definition
phase.cpp:23
mim::RWPhase::new_world
World & new_world()
Create new Defs into this.
Definition
phase.h:368
mim::Rewriter::rewrite
virtual const Def * rewrite(const Def *)
Definition
rewrite.cpp:56
mim::plug::tensor::phase::Lower::rewrite_imm_App
const Def * rewrite_imm_App(const App *) final
Definition
lower.cpp:33
def.h
lam.h
lower.h
mim::plug::math::tri::h
@ h
Definition
autogen.h:169
mim::plug::regex::cls::w
@ w
Definition
autogen.h:63
mim::plug::tensor::phase
Definition
fuse.h:5
mim::plug::tensor::slice_impl
slice_impl
Definition
autogen.h:98
mim::plug::tensor::bmm_impl
bmm_impl
Definition
autogen.h:199
mim::plug::tensor::flip_impl
flip_impl
Definition
autogen.h:113
mim::plug::tensor::conv_impl
conv_impl
Definition
autogen.h:143
mim::plug::tensor::select_impl
select_impl
Definition
autogen.h:306
mim::plug::tensor::product_2d_impl
product_2d_impl
Definition
autogen.h:185
mim::plug::tensor::transpose_2d_impl
transpose_2d_impl
Definition
autogen.h:227
mim::plug::tensor::unary_impl
unary_impl
Definition
autogen.h:278
mim::plug::tensor::repeat_impl
repeat_impl
Definition
autogen.h:68
mim::plug::tensor::map_impl
map_impl
Definition
autogen.h:264
mim::plug::tensor::transpose_impl
transpose_impl
Definition
autogen.h:213
mim::plug::tensor::binary_impl
binary_impl
Definition
autogen.h:292
mim::plug::tensor::dot_product_impl
dot_product_impl
Definition
autogen.h:171
mim::plug::tensor::reshape_impl
reshape_impl
Definition
autogen.h:83
mim::plug::tensor::broadcast_in_dim_impl
broadcast_in_dim_impl
Definition
autogen.h:250
mim::plug::tensor::pool_impl
pool_impl
Definition
autogen.h:157
mim::plug::tuple::head
head
Definition
autogen.h:43
mim::DefVec
Vector< const Def * > DefVec
Definition
def.h:79
mim::Node::App
@ App
Definition
def.h:109
tensor.h
src
mim
plug
tensor
phase
lower.cpp
Generated by
1.16.1