16 auto dom = outer->
type()->as<
Pi>()->dom();
17 auto codom = inner->
type()->as<
Pi>()->codom();
18 auto lam = w.mut_lam(dom, codom)->set(
"fused_map");
19 lam->set(
true, w.app(inner, w.app(outer, lam->
var())));
44const Def* Fuse::fuse_map_reduce(
const App* app) {
47 auto [nis, meta, shapes, TisRisSis, comb_init, map_out, maps] = outer_callee->
uncurry_args<7>();
49 auto [To, Ro, Rr] = meta->
projs<3>();
50 auto [comb, init] = comb_init->
projs<2>();
51 auto [Tis, Ris, Sis] = TisRisSis->
projs<3>();
54 DLOG(
"considering map_reduce for fusion:");
55 DLOG(
" comb = {} : {}", comb, comb->type());
56 DLOG(
" init = {} : {}", init, init->type());
57 DLOG(
" Tis = {} : {}", Tis, Tis->type());
58 DLOG(
" Ris = {} : {}", Ris, Ris->type());
59 DLOG(
" Sis = {} : {}", Sis, Sis->type());
60 DLOG(
" To = {} : {}", To, To->type());
61 DLOG(
" Ro = {} : {}", Ro, Ro->type());
62 DLOG(
" nis = {} : {}", nis, nis->
type());
63 DLOG(
" is = {} : {}", is, is->type());
68 if (!nis_lit)
return nullptr;
69 auto nis_nat = *nis_lit;
73 const Def* comb =
nullptr;
74 const Def* init =
nullptr;
75 const Def* Tis =
nullptr;
76 const Def* Ris =
nullptr;
77 const Def* Sis =
nullptr;
78 const Def* To =
nullptr;
79 const Def* maps =
nullptr;
81 const Def* is =
nullptr;
85 bool any_fusible =
false;
87 for (
u64 k = 0; k < nis_nat; ++k) {
88 auto input_k = is->proj(nis_nat, k);
92 auto [inner_nis, inner_meta, inner_shapes, inner_TisRisSis, inner_comb_init, inner_map_out, inner_maps,
94 = inner->uncurry_args<8>();
95 auto [inner_To, inner_Ro, inner_Rr] = inner_meta->projs<3>();
96 auto [inner_So, inner_Sr] = inner_shapes->projs<2>();
97 auto [inner_comb, inner_init] = inner_comb_init->projs<2>();
98 auto [inner_Tis, inner_Ris, inner_Sis] = inner_TisRisSis->projs<3>();
101 if (!inner_nis_nat)
continue;
109 if (!inner_rr || *inner_rr != 0)
continue;
110 if (inner_Sr != inner_So)
continue;
111 auto id_lam = inner_map_out->isa_mut<
Lam>();
112 if (!id_lam || !id_lam->is_set() || id_lam->body() != id_lam->var())
continue;
114 auto& info = infos[k];
116 info.comb = inner_comb;
117 info.init = inner_init;
118 info.Tis = inner_Tis;
119 info.Ris = inner_Ris;
120 info.Sis = inner_Sis;
122 info.maps = inner_maps;
123 info.nis = *inner_nis_nat;
128 if (!any_fusible)
return nullptr;
135 for (
u64 i = 0; i < nis_nat; ++i) {
136 new_pos[i] = new_nis_nat;
137 new_nis_nat += infos[i].fusible ? infos[i].nis : 1;
140 DefVec new_Tis_vec(new_nis_nat);
141 DefVec new_Ris_vec(new_nis_nat);
142 DefVec new_Sis_vec(new_nis_nat);
143 DefVec new_maps_vec(new_nis_nat);
144 DefVec new_is_vec(new_nis_nat);
146 for (u64 i = 0; i < nis_nat; ++i) {
147 if (infos[i].fusible) {
148 const auto& info = infos[i];
149 auto outer_map_i = maps->proj(nis_nat, i);
150 for (u64 l = 0; l < info.nis; ++l) {
151 auto pos = new_pos[i] + l;
152 new_Tis_vec[pos] = info.Tis->proj(info.nis, l);
153 new_Ris_vec[pos] = info.Ris->proj(info.nis, l);
154 new_Sis_vec[pos] = info.Sis->proj(info.nis, l);
155 new_is_vec[pos] = info.is->proj(info.nis, l);
158 new_maps_vec[pos] =
compose_map(w, info.maps->proj(info.nis, l), outer_map_i);
161 auto pos = new_pos[i];
162 new_Tis_vec[pos] = Tis->proj(nis_nat, i);
163 new_Ris_vec[pos] = Ris->proj(nis_nat, i);
164 new_Sis_vec[pos] = Sis->proj(nis_nat, i);
165 new_maps_vec[pos] = maps->proj(nis_nat, i);
166 new_is_vec[pos] = is->proj(nis_nat, i);
170 auto new_Tis =
w.tuple(new_Tis_vec);
171 auto new_Ris =
w.tuple(new_Ris_vec);
172 auto new_Sis =
w.tuple(new_Sis_vec);
173 auto new_maps =
w.tuple(new_maps_vec);
174 auto new_is =
w.tuple(new_is_vec);
176 auto new_nis_def =
w.lit_nat(new_nis_nat);
193 auto inputs_sigma =
w.sigma(new_Tis_vec);
194 auto data_sigma =
w.sigma({To, inputs_sigma});
195 auto ret_cn_type =
w.cn(To);
196 auto new_comb =
w.mut_con({data_sigma, ret_cn_type})->
set(
"fused_comb");
197 auto new_data = new_comb->var(0);
198 auto new_ret = new_comb->var(1);
199 auto new_acc = new_data->proj(2, 0);
200 auto new_in = new_data->proj(2, 1);
203 for (
u64 i = 0; i < nis_nat; ++i)
204 if (infos[i].fusible) fused_indices.emplace_back(i);
207 Vector<const Def*> inner_values(fused_indices.size());
208 for (
size_t r = 0;
r < fused_indices.size(); ++
r) {
209 auto new_inner_To = infos[fused_indices[
r]].To;
210 inner_rets[
r] =
w.mut_con(new_inner_To)->set(
"inner_ret");
211 inner_values[
r] = inner_rets[
r]->var(0);
215 DefVec outer_inputs_vec(nis_nat);
218 for (
u64 i = 0; i < nis_nat; ++i)
219 if (infos[i].fusible)
220 outer_inputs_vec[i] = inner_values[
r++];
222 outer_inputs_vec[i] = new_in->proj(new_nis_nat, new_pos[i]);
226 for (
size_t r = 0;
r < fused_indices.size(); ++
r) {
227 auto k = fused_indices[
r];
228 auto new_inner_comb = infos[k].comb;
229 auto new_inner_init = infos[k].init;
231 DefVec inner_inputs_vec(infos[k].nis);
232 for (
u64 l = 0;
l < infos[k].nis; ++
l)
233 inner_inputs_vec[l] = new_in->proj(new_nis_nat, new_pos[k] + l);
235 Lam* caller = (
r == 0) ? new_comb : inner_rets[
r - 1];
236 caller->app(
true, new_inner_comb, {
w.tuple({new_inner_init,
w.tuple(inner_inputs_vec)}), inner_rets[r]});
240 inner_rets.back()->app(
true, comb, {
w.tuple({new_acc,
w.tuple(outer_inputs_vec)}), new_ret});
244 mr =
w.app(mr, new_nis_def);
245 mr =
w.app(mr, meta);
246 mr =
w.app(mr, shapes);
247 mr =
w.app(mr, {new_Tis, new_Ris, new_Sis});
248 mr =
w.app(mr, {new_comb,
init});
249 mr =
w.app(mr, map_out);
250 mr =
w.app(mr, new_maps);
251 mr =
w.app(mr, new_is);
258 if (
auto res = fuse_map_reduce(mr)) {
259 DLOG(
"Fused map_reduce at {} into a new map_reduce {}", app, res);
263 return RWPhase::rewrite_imm_App(app);
const Def * callee() const
static auto uncurry_args(const Def *def)
static auto isa(const Def *def)
const Def * var(nat_t a, nat_t i) noexcept
auto projs(F f) const
Splits this Def via Def::projections into an Array (if A == std::dynamic_extent) or std::array (other...
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
static std::optional< T > isa(const Def *def)
A dependent function type.
World & new_world()
Create new Defs into this.
virtual const Def * rewrite(const Def *)
This is a thin wrapper for absl::InlinedVector<T, N, A> which is a drop-in replacement for std::vecto...
The World represents the whole program and manages creation of MimIR nodes (Defs).
const Def * rewrite_imm_App(const App *) final
#define DLOG(...)
Vaporizes to nothingness in Debug build.
static const Def * compose_map(World &w, const Def *inner, const Def *outer)
inner ∘ outer: feeds the outer op's read coordinates for one input into the inner op's access map.
Vector< const Def * > DefVec
Vector(I, I, A=A()) -> Vector< typename std::iterator_traits< I >::value_type, Default_Inlined_Size< typename std::iterator_traits< I >::value_type >, A >