28 if (e == var)
return 0;
32 auto [x, c] = semiop->args<2>();
39 auto [a, b] = op->args<2>();
42 if (a_const == b_const)
return {};
58 auto lam =
map->isa_mut<
Lam>();
59 if (!lam || !lam->is_set())
return false;
60 auto var = lam->
var();
61 auto body = lam->body();
62 if (body == var)
return true;
67 auto used = fe::Bitset();
68 auto mark = [&](
const Def* elem) {
70 if (!i || *i >= *n)
return false;
76 for (
auto elem :
tuple->ops())
77 if (!mark(elem))
return false;
78 }
else if (
auto pack = body->isa<
Pack>()) {
80 if (!pack->is_set() || !mark(pack->body()))
return false;
81 }
else if (!mark(body)) {
85 return used.count() == *n;
90 auto dom = outer->
type()->as<
Pi>()->dom();
91 auto codom = inner->
type()->as<
Pi>()->codom();
92 auto lam = w.mut_lam(dom, codom)->set(
"fused_map");
93 lam->set(
true, w.app(inner, w.app(outer, lam->
var())));
103 T.emplace_back(t),
R.emplace_back(r),
S.emplace_back(s),
maps.emplace_back(m),
is.emplace_back(v);
115 if (!pr->src->type()->isa<
Arr>())
return {};
120 auto [T, r] = bc->callee()->as<
App>()->args<2>();
121 auto [s_in, s_out, input] = bc->arg()->projs<3>();
125 if (!input->type()->isa<
Arr>())
return {};
129 auto vec_ty = slot_map->
type()->as<
Pi>()->codom();
130 auto lam = w.mut_lam(vec_ty, vec_ty)->set(
"bcast_map");
132 for (
u64 d = 0; d < *r_l; ++d) {
133 auto in_d = s_in->proj(*r_l, d);
134 auto o_d = lam->var(*r_l, d);
135 if (in_d == s_out->proj(*r_l, d))
142 lam->set(
true, w.tuple(elems));
144 return PureRead{input, lam, T, r, s_in};
173const Def* Fuse::fuse_map_reduce(
const App* app) {
176 auto [nis_nps, meta, shapes, in_tys, comb_init, map_out, maps_all] = outer_callee->
uncurry_args<7>();
178 auto [nis, nps] = nis_nps->
projs<2>();
179 auto [To, Tp, Ro, Rn, TSched] = meta->
projs<5>();
180 auto [comb, init, post] = comb_init->
projs<3>();
181 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->
projs<6>();
182 auto [maps, post_maps] = maps_all->projs<2>();
184 auto [is, post_is] = is_all->projs<2>();
186 log().d(
"consider for fusion: comb = {}, init = {}, To = {}, Ro = {}", comb, init, To, Ro);
187 log().d(
" inputs: nis = {}, Tis = {}, Ris = {}, Sis = {}, is = {}", nis, Tis, Ris, Sis, is);
192 if (!nis_lit)
return nullptr;
193 auto nis_nat = *nis_lit;
207 fe::Vector<std::optional<InnerInfo>> infos(nis_nat);
209 for (
u64 k = 0; k < nis_nat; ++k) {
211 if (!inner)
continue;
213 auto [inner_nis_nps, inner_meta, inner_shapes, inner_in_tys, inner_comb_init, inner_map_out, inner_maps_all,
215 = inner->uncurry_args<8>();
216 auto [inner_nis, inner_nps] = inner_nis_nps->projs<2>();
217 auto [inner_To, inner_Tp, inner_Ro, inner_Rn, inner_TSched] = inner_meta->projs<5>();
218 auto [inner_So, inner_Sr, inner_sched] = inner_shapes->projs<3>();
219 auto [inner_comb, inner_init, inner_post] = inner_comb_init->projs<3>();
220 auto [inner_Tis, inner_Ris, inner_Sis, inner_Tps, inner_Rps, inner_Sps] = inner_in_tys->projs<6>();
221 auto [inner_maps, inner_post_maps] = inner_maps_all->projs<2>();
222 auto [inner_is, inner_post_is] = inner_is_all->projs<2>();
225 if (!inner_nis_nat)
continue;
230 if (!inner_nps_nat || *inner_nps_nat != 0)
continue;
239 if (!inner_ro || !inner_rn || *inner_ro != *inner_rn)
continue;
240 if (inner_Sr != inner_So)
continue;
241 auto id_lam = inner_map_out->isa_mut<
Lam>();
242 if (!id_lam || !id_lam->is_set() || id_lam->body() != id_lam->var())
continue;
255 infos[k] = InnerInfo{.comb = inner_comb,
262 .nis = *inner_nis_nat,
266 if (std::ranges::none_of(infos, [](
const auto& info) {
return info.has_value(); }))
return nullptr;
271 fe::Vector<u64> new_pos(nis_nat);
273 for (
u64 i = 0; i < nis_nat; ++i) {
274 new_pos[i] = new_nis_nat;
275 new_nis_nat += infos[i] ? infos[i]->nis : 1;
278 DefVec new_Tis_vec(new_nis_nat);
279 DefVec new_Ris_vec(new_nis_nat);
280 DefVec new_Sis_vec(new_nis_nat);
281 DefVec new_maps_vec(new_nis_nat);
282 DefVec new_is_vec(new_nis_nat);
284 for (
u64 i = 0; i < nis_nat; ++i) {
285 if (
auto& info = infos[i]) {
286 auto outer_map_i = maps->proj(nis_nat, i);
289 bool shared = !sole_consumer(is->proj(nis_nat, i));
291 auto pos = new_pos[i] +
l;
292 new_Tis_vec[pos] =
info->Tis->proj(
info->nis, l);
293 new_Ris_vec[pos] =
info->Ris->proj(
info->nis, l);
294 new_Sis_vec[pos] =
info->Sis->proj(
info->nis, l);
295 new_is_vec[pos] =
info->is->proj(
info->nis, l);
296 if (shared) shared_.insert(new_is_vec[pos]);
302 auto pos = new_pos[i];
303 new_Tis_vec[pos] = Tis->proj(nis_nat, i);
304 new_Ris_vec[pos] = Ris->proj(nis_nat, i);
305 new_Sis_vec[pos] = Sis->proj(nis_nat, i);
306 new_maps_vec[pos] = maps->proj(nis_nat, i);
307 new_is_vec[pos] = is->proj(nis_nat, i);
311 auto new_Tis =
w.tuple(new_Tis_vec);
312 auto new_Ris =
w.tuple(new_Ris_vec);
313 auto new_Sis =
w.tuple(new_Sis_vec);
314 auto new_maps =
w.tuple(new_maps_vec);
315 auto new_is =
w.tuple(new_is_vec);
317 auto new_nis_def =
w.lit_nat(new_nis_nat);
334 auto inputs_sigma =
w.sigma(new_Tis_vec);
335 auto data_sigma =
w.sigma({To, inputs_sigma});
336 auto ret_cn_type =
w.cn(To);
337 auto new_comb =
w.mut_con({data_sigma, ret_cn_type})->
set(
"fused_comb");
338 auto [new_data, new_ret] = new_comb->vars<2>();
339 auto [new_acc, new_in] = new_data->projs<2>();
341 fe::Vector<u64> fused_indices;
342 for (
u64 i = 0; i < nis_nat; ++i)
343 if (infos[i]) fused_indices.emplace_back(i);
345 fe::Vector<Lam*> inner_rets(fused_indices.size(),
346 [&](
size_t r) { return w.mut_con(infos[fused_indices[r]]->To)->set(
"inner_ret"); });
349 DefVec outer_inputs_vec(nis_nat);
350 for (
u64 i = 0, r = 0; i < nis_nat; ++i)
351 outer_inputs_vec[i] = infos[i] ? inner_rets[r++]->var() : new_in->proj(new_nis_nat, new_pos[i]);
354 for (
size_t r = 0;
r < fused_indices.size(); ++
r) {
355 const auto&
info = *infos[fused_indices[
r]];
357 [&](
size_t l) { return new_in->proj(new_nis_nat, new_pos[fused_indices[r]] + l); });
358 Lam* caller =
r == 0 ? new_comb : inner_rets[
r - 1];
359 caller->app(
true,
info.comb, {w.tuple({info.init, w.tuple(inner_inputs_vec)}), inner_rets[
r]});
363 inner_rets.back()->app(
true, comb, {
w.tuple({new_acc,
w.tuple(outer_inputs_vec)}), new_ret});
368 mr =
w.app(mr, {new_nis_def, nps});
369 mr =
w.app(mr, meta);
370 mr =
w.app(mr, shapes);
371 mr =
w.app(mr, {new_Tis, new_Ris, new_Sis, Tps, Rps, Sps});
372 mr =
w.app(mr, {new_comb,
init, post});
373 mr =
w.app(mr, map_out);
374 mr =
w.app(mr, {new_maps, post_maps});
375 mr =
w.app(mr, {new_is, post_is});
408 auto pi =
map->type()->isa<
Pi>();
409 if (!pi)
return false;
410 auto expected = w.app(w.app(w.annex<
tensor::reshape_map>(), {r_in, r_out}), {s_in, s_out});
411 auto epi = expected->type()->isa<
Pi>();
412 if (!epi || epi->dom() != pi->dom() || epi->codom() != pi->codom())
return false;
413 auto probe = w.mut_lam(pi->dom(), pi->codom());
414 return w.app(
map, probe->var()) == w.app(expected, probe->var());
420bool Fuse::sole_consumer(
const Def* d)
const {
421 if (shared_.contains(d))
return false;
422 auto old_it = new2old_.find(d);
423 if (old_it == new2old_.end())
return false;
424 auto cnt = mr_consumers_.find(old_it->second);
425 return cnt != mr_consumers_.end() && cnt->second == 1;
428const Def* Fuse::fuse_epilogue(
const App* callee,
const Def* arg) {
429 auto [nis_nps, meta, shapes, in_tys, comb_init, map_out, maps_all] = callee->
uncurry_args<7>();
431 auto [nis, nps] = nis_nps->
projs<2>();
432 auto [To, Tp, Ro, Rn, TSched] = meta->projs<5>();
433 auto [So, Sr, sched] = shapes->projs<3>();
434 auto [comb, init, post] = comb_init->projs<3>();
435 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->projs<6>();
436 auto [maps, post_maps] = maps_all->projs<2>();
438 auto& w = new_world();
444 if (!nis_lit || !nps_lit || !ro_lit || !rn_lit || *ro_lit != *rn_lit)
return nullptr;
445 auto nis_nat = *nis_lit;
446 auto nps_nat = *nps_lit;
447 if (nis_nat == 0)
return nullptr;
448 if (Sr != So)
return nullptr;
450 auto id_out = map_out->isa_mut<
Lam>();
451 if (!id_out || !id_out->is_set() || id_out->body() != id_out->var())
return nullptr;
457 auto [is, ps] = arg->
projs<2>();
463 const Def* inner_def =
nullptr;
465 for (
u64 k = 0; k < nis_nat && !inner_def; ++k) {
466 auto cand = is->proj(nis_nat, k);
468 if (!sole_consumer(cand))
continue;
470 auto m = maps->proj(nis_nat, k);
471 auto id_in = m->isa_mut<
Lam>();
472 if (id_in && id_in->is_set() && id_in->body() == id_in->var() && Sis->proj(nis_nat, k) == So) {
475 }
else if (Sis->proj(nis_nat, k) != So
476 && is_unpack_read(w, m, Ris->proj(nis_nat, k), Sis->proj(nis_nat, k), Ro, So)) {
482 if (!inner_def)
return nullptr;
488 const Def* to_cells =
nullptr;
491 =
w.app(
w.app(
w.annex<
tensor::reshape_map>(), {Ro, Ris->proj(nis_nat, k0)}), {So, Sis->proj(nis_nat, k0)});
494 auto [i_nis_nps, i_meta, i_shapes, i_in_tys, i_comb_init, i_map_out, i_maps_all, i_is_all]
495 = inner->uncurry_args<8>();
496 auto [i_nis, i_nps] = i_nis_nps->projs<2>();
497 auto [i_To, i_Tp, i_Ro, i_Rn, i_TSched] = i_meta->projs<5>();
498 auto [i_Tis, i_Ris, i_Sis, i_Tps, i_Rps, i_Sps] = i_in_tys->projs<6>();
499 auto [i_comb, i_init, i_post] = i_comb_init->projs<3>();
500 auto [i_maps, i_post_maps] = i_maps_all->projs<2>();
501 auto [i_is, i_ps] = i_is_all->projs<2>();
504 if (!i_nps_lit)
return nullptr;
505 auto i_nps_nat = *i_nps_lit;
507 log().d(
"fuse trailing map {} {} into the epilogue of {}", callee, arg, inner_def);
512 auto new_nps = i_nps_nat + (nis_nat - 1) + nps_nat;
513 auto to_cell = [&](
const Def* m) {
return to_cells ?
compose_map(w, m, to_cells) : m; };
515 for (
u64 j = 0; j < i_nps_nat; ++j)
516 eps.push(i_Tps->proj(i_nps_nat, j), i_Rps->proj(i_nps_nat, j), i_Sps->proj(i_nps_nat, j),
517 i_post_maps->proj(i_nps_nat, j), i_ps->proj(i_nps_nat, j));
518 for (
u64 i = 0; i < nis_nat; ++i)
520 eps.push(Tis->proj(nis_nat, i), Ris->proj(nis_nat, i), Sis->proj(nis_nat, i),
521 to_cell(maps->proj(nis_nat, i)), is->proj(nis_nat, i));
522 for (
u64 j = 0; j < nps_nat; ++j)
523 eps.push(Tps->proj(nps_nat, j), Rps->proj(nps_nat, j), Sps->proj(nps_nat, j),
524 to_cell(post_maps->proj(nps_nat, j)), ps->proj(nps_nat, j));
528 auto fused_post =
w.mut_con({
w.sigma({i_To,
w.sigma(eps.T)}),
w.cn(Tp)})->
set(
"fused_post");
529 auto after_ip =
w.mut_con(i_Tp)->set(
"afterInnerPost");
530 auto after_comb =
w.mut_con(To)->set(
"afterComb");
531 auto [fused_data, fused_ret] = fused_post->vars<2>();
532 auto [x, extras] = fused_data->projs<2>();
534 DefVec i_extras(i_nps_nat, [&](
size_t j) {
return extras->proj(new_nps, j); });
535 fused_post->app(
true, i_post, {
w.tuple({x,
w.tuple(i_extras)}), after_ip});
537 DefVec comb_inputs(nis_nat);
538 for (
u64 i = 0, r = 0; i < nis_nat; ++i)
539 comb_inputs[i] = i == k0 ? after_ip->var() : extras->proj(new_nps, i_nps_nat + r++);
540 after_ip->app(
true, comb, {
w.tuple({
init,
w.tuple(comb_inputs)}), after_comb});
542 DefVec o_extras(nps_nat, [&](
size_t j) {
return extras->proj(new_nps, i_nps_nat + (nis_nat - 1) + j); });
543 after_comb->app(
true, post, {
w.tuple({after_comb->var(),
w.tuple(o_extras)}), fused_ret});
548 mr =
w.app(mr, {i_nis,
w.lit_nat(new_nps)});
549 mr =
w.app(mr, {i_To, Tp, i_Ro, i_Rn, i_TSched});
550 mr =
w.app(mr, i_shapes);
551 mr =
w.app(mr, {i_Tis, i_Ris, i_Sis,
w.tuple(eps.T),
w.tuple(eps.R),
w.tuple(eps.S)});
552 mr =
w.app(mr, {i_comb, i_init, fused_post});
553 mr =
w.app(mr, i_map_out);
554 mr =
w.app(mr, {i_maps,
w.tuple(eps.maps)});
555 mr =
w.app(mr, {i_is,
w.tuple(eps.is)});
563 impl =
w.app(impl, {Tp, Ris->proj(nis_nat, k0), Ro});
564 impl =
w.app(impl, Sis->proj(nis_nat, k0));
565 impl =
w.app(impl, So);
566 mr =
w.app(impl, mr);
578const Def* Fuse::fuse_read_through(
const App* callee,
const Def* arg) {
579 auto [nis_nps, meta, shapes, in_tys, comb_init, map_out, maps_all] = callee->uncurry_args<7>();
581 auto [nis, nps] = nis_nps->
projs<2>();
582 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->projs<6>();
583 auto [maps, post_maps] = maps_all->projs<2>();
587 if (!nis_lit || !nps_lit)
return nullptr;
588 auto nis_nat = *nis_lit;
589 auto nps_nat = *nps_lit;
591 auto&
w = new_world();
593 auto [is, ps] = arg->projs<2>();
595 bool changed =
false;
596 auto rewire = [&](
u64 n,
const Def* Ts,
const Def* Rs,
const Def* Ss,
const Def* ms,
const Def* vs) {
597 auto slots = Slots();
598 for (
u64 i = 0; i < n; ++i) {
599 auto T = Ts->proj(n, i),
R = Rs->proj(n, i),
S = Ss->proj(n, i);
600 auto m = ms->proj(n, i), v = vs->proj(n, i);
602 log().d(
"read input {} of {} through {}", i, callee, v);
605 if (!sole_consumer(v)) shared_.insert(
rt->src);
609 slots.push(T, R, S, m, v);
614 auto ins = rewire(nis_nat, Tis, Ris, Sis, maps, is);
615 auto eps = rewire(nps_nat, Tps, Rps, Sps, post_maps, ps);
616 if (!changed)
return nullptr;
619 mr =
w.app(mr, nis_nps);
620 mr =
w.app(mr, meta);
621 mr =
w.app(mr, shapes);
622 mr =
w.app(mr, {
w.tuple(ins.T),
w.tuple(ins.R),
w.tuple(ins.S),
w.tuple(eps.T),
w.tuple(eps.R),
w.tuple(eps.S)});
623 mr =
w.app(mr, comb_init);
624 mr =
w.app(mr, map_out);
625 mr =
w.app(mr, {
w.tuple(ins.maps),
w.tuple(eps.maps)});
626 mr =
w.app(mr, {
w.tuple(ins.is),
w.tuple(eps.is)});
645 const App* cur =
nullptr;
646 if (
auto res = fuse_map_reduce(app)) {
647 log().d(
"fused map_reduce {} → {}", app, res);
648 cur = res->as<
App>();
657 for (
bool progress =
true; progress;) {
659 const Def* res = fuse_read_through(callee, arg);
661 res = fuse_epilogue(callee, arg);
662 if (res)
log().d(
"fused trailing map {} into its producer's epilogue → {}", app, res);
665 cur = res->as<
App>();
671 auto result = cur ? cur : RWPhase::rewrite_imm_App(app);
675 new2old_[result] = app;
680 if (
auto pr =
is_pure_read(result); pr && is_mr(pr->src) && !new2old_.contains(pr->src))
681 new2old_[pr->src] = app;
685 return RWPhase::rewrite_imm_App(app);
const Def * callee() const
static auto uncurry_args(const Def *def)
A (possibly paramterized) Array.
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 (possibly paramterized) Tuple.
const fe::Log & log() const
A dependent function type.
World & new_world()
Create new Defs into this.
void start() override
RWBase::start() and then swaps the two worlds.
World & old_world()
Get old Defs from here.
virtual const Def * rewrite(const Def *)
Data constructor for a Sigma.
The World represents the whole program and manages creation of MimIR nodes (Defs).
void start() override
Actual entry.
const Def * rewrite_imm_App(const App *) final
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.
static std::optional< PureRead > read_through(World &w, const Def *value, const Def *slot_map)
If value is a pure re-indexed read — a copy-combiner map_reduce (reshape/transpose/slice/ flip/repeat...
static bool is_unpack_read(World &w, const Def *map, const Def *r_in, const Def *s_in, const Def *r_out, const Def *s_out)
Is map the row-major reshape read tensor.reshape_map (s_in, s_out) — the map that reads a PACKED prod...
static std::optional< u64 > injective_coord(const Def *var, const Def *e)
If e reads coordinate var#i injectively, returns i: a plain extract, possibly strided (affine....
static bool reads_injectively(const Def *map)
Checks that map provably reads through every loop index of its domain: its body is the identity,...
bool is_identity_post(const Def *post)
Is post the (rebuilt) CPS identity tensor.id, i.e.
std::optional< PureRead > is_pure_read(const Def *value)
If value is a pure re-indexed read — a copy-combiner map_reduce without reduction loops that writes i...
bool is_copy_comb(const Def *comb)
Recognizes the (rebuilt) tensor_copy combiner (acc, ys) ↦ ys#0: the result is exactly the single inpu...
DefMap< u64 > count_consumers(const World &world, Pred pred)
Counts the consumers of every def of world matched by pred.
A pure re-indexed read: the source tensor, the access map into it (over the read's output coordinates...
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec
The five parallel per-slot lists of a map_reduce_post input group: element type, rank,...
void push(const Def *t, const Def *r, const Def *s, const Def *m, const Def *v)