21const Def* fold_index(
const Def* shape,
const Def* idx) {
23 auto r =
shape->num_projs();
26 for (
size_t i = 0; i !=
r; ++i)
30 out.push_back(
idx->proj(r, i));
34 if (!dropped)
return idx;
39std::pair<Lam*, const Def*> counting_for(
const Def* bound,
const Def* acc,
const Def* exit, Sym name) {
40 auto&
w = bound->world();
41 auto acc_ty = acc->type();
42 auto body =
w.mut_con({
w.type_i64(), acc_ty,
w.cn(acc_ty)})->
set(name);
43 auto for_loop =
w.call<
affine::For>(body, exit,
Defs{
w.lit_i64(0), bound,
w.lit_i64(1), acc});
44 return {body, for_loop};
52template<
class Compute>
53const Def* build_pointwise(World& w,
59 const std::string& name,
61 auto mem_ty =
w.call<
mem::M>(0);
62 auto fun =
w.mut_fun(
w.sigma({mem_ty, op_ins->type()}), result_ty)->set(name);
64 auto [fun_mem, ins] = fun->var(0_n)->projs<2>();
65 auto cont = fun->var(1);
69 const Def* acc =
w.tuple({a_mem, out_buf});
70 auto current_mut = fun;
74 for (
u64 d = 0;
d < rn; ++
d) {
76 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(name +
"_" + std::to_string(d)));
77 auto [iter, new_acc, yield] = body->vars<3>();
79 iters.push_back(iter);
81 current_mut->set(
true, for_call);
84 auto [loop_mem, loop_buf] = acc->projs<2>();
86 std::pair<const Def*, const Def*> el = compute(iters, ins, loop_mem);
87 auto [el_mem, elem] = el;
90 for (
u64 d = 0;
d < rn; ++
d)
93 =
buffer::op_write(obr, obs, obT, el_mem, loop_buf, fold_index(s_out,
w.tuple(wcoords)), elem)->
projs<2>();
94 current_mut->app(
true, cont,
w.tuple({wr_mem, loop_buf}));
109 return RWPhase::rewrite_imm_App(app);
112const Def* LowerMapReduce::lower_buffer_lit(
const App* app) {
119 if (!rn)
return RWPhase::rewrite_imm_App(app);
125 return build_pointwise(
126 w, result_ty,
mem, val, s_out, *rn,
"constant_fill",
127 [](
const DefVec&,
const Def* ins,
const Def* m) -> std::pair<const Def*, const Def*> {
return {m, ins}; });
130const Def* LowerMapReduce::lower_map_reduce_post(
const App* app) {
134 auto [nis_nps, meta, shapes, in_tys, comb_init, acc_out, accs_all] = c->uncurry_args<7>();
135 auto [nis, nps] = nis_nps->
projs<2>();
136 auto [To, Tp, Ro, Rn, TSched] = meta->
projs<5>();
137 auto [So, Sr, sched] = shapes->
projs<3>();
138 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->
projs<6>();
139 auto [comb, init, post] = comb_init->
projs<3>();
140 auto [accs, post_accs] = accs_all->projs<2>();
149 if (!nis_l || !nps_l || !ro_l || !rn_l || *rn_l < *ro_l) {
150 log().w(
"rank counts (nis/nps/Ro/Rn) of {} are not known at lowering time", app);
151 return RWPhase::rewrite_imm_App(app);
153 auto nis_nat = *nis_l;
154 auto nps_nat = *nps_l;
155 auto ro = *ro_l, rr = *rn_l - *ro_l;
157 auto n = w.lit_nat(nloops);
161 auto affine_map = [&](
const Def* f,
const Def* m,
const Def* nn,
const Def* sin,
const Def* sout,
const Def* idxs,
163 auto a = w.app(w.annex<
affine::map>(), w.tuple({m, nn}));
164 a = w.app(a, w.tuple({sin, sout}));
167 a = w.app(a, w.lit_nat_0());
168 return w.app(a, mem)->projs<2>();
171 auto mem_ty =
w.call<
mem::M>(0);
174 auto fun =
w.mut_fun(
w.sigma({mem_ty, op_is->type(), op_post_is->type()}), result_ty)->set(
"mapRedAff");
176 auto [fun_mem, new_inputs, new_post_is] = fun->var(0_n)->projs<3>();
177 auto cont = fun->var(1);
186 auto i32 =
w.type_i32();
192 auto nest_args =
w.tuple({Ro,
w.lit_nat(rr), Sr, To, result_ty->proj(1)});
198 auto sched_dom = nest->type()->as<
Pi>()->dom();
203 auto apply_cps = [&](
Lam* mut,
const Def*
f,
DefVec parts,
const Def* k) {
204 auto dom =
f->type()->as<
Pi>()->dom();
205 if (dom->num_projs() == parts.size() + 1) {
206 parts.emplace_back(k);
207 mut->app(
true, f,
w.tuple(parts));
209 mut->app(
true, f,
w.tuple({w.tuple(parts), k}));
217 if (ndom == pos + cnt + 1)
218 for (
u64 d = 0;
d < cnt; ++
d)
219 out[d] =
l->var(ndom, pos + d);
221 for (
u64 d = 0;
d < cnt; ++
d)
222 out[d] =
l->var(ndom, pos)->proj(cnt, d);
227 auto cdom = sched_dom->proj(3, 1)->as<
Pi>()->dom();
228 auto cn = cdom->num_projs();
229 auto cell =
w.mut_con(cdom)->set(
"cell");
231 auto cm = cell->var(cn, 0);
232 auto cacc = cell->var(cn, 1);
233 auto ck = cell->var(cn, cn - 1);
234 auto civs = load_ivs(cell, cn, 2, nloops);
236 for (
u64 d = 0;
d < nloops; ++
d)
238 auto iters =
w.tuple(iters_v);
240 DefVec input_elems(nis_nat);
241 for (
u64 i = 0; i < nis_nat; ++i) {
242 auto in_buf = new_inputs->proj(nis_nat, i);
243 auto [mc_mem, coords]
244 = affine_map(accs->proj(nis_nat, i), Ris->proj(nis_nat, i), n, Sr, Sis->proj(nis_nat, i), iters, cur);
247 auto [rd_mem, rd_val]
248 =
buffer::op_read(ir, is_, iT, cur, in_buf, fold_index(Sis->proj(nis_nat, i), coords))->
projs<2>();
250 input_elems[i] = rd_val;
252 apply_cps(cell, comb, {cur, cacc,
w.tuple(input_elems)}, ck);
259 auto wdom = sched_dom->proj(3, 2)->as<
Pi>()->dom();
260 auto wn = wdom->num_projs();
261 auto wb =
w.mut_con(wdom)->set(
"wb");
263 auto wm = wb->var(wn, 0);
264 auto wu = wb->var(wn, 1);
265 auto wv = wb->var(wn, 2);
266 auto wk = wb->var(wn, wn - 1);
267 auto wovs = load_ivs(wb, wn, 3, nloops);
269 for (
u64 i = 0; i < ro; ++i)
270 wb_iters[i] =
w.call(
core::conv::u, Sr->proj(nloops, i), wovs[i]);
271 for (
u64 j = 0; j < rr; ++j)
272 wb_iters[ro + j] =
w.call(
core::conv::u, Sr->proj(nloops, ro + j),
w.lit(i32, 0));
273 auto [wc_mem, write_coords] = affine_map(acc_out, Ro, n, Sr, So,
w.tuple(wb_iters), wm);
276 DefVec post_elems(nps_nat);
277 for (
u64 j = 0; j < nps_nat; ++j) {
278 auto sps_j = Sps->proj(nps_nat, j);
279 auto [pc_mem, pcoords]
280 = affine_map(post_accs->proj(nps_nat, j), Rps->proj(nps_nat, j), Ro, So, sps_j, write_coords, pcur);
282 auto p_buf = new_post_is->proj(nps_nat, j);
284 auto [prd_mem, p_val] =
buffer::op_read(pr, ps_, pT, pcur, p_buf, fold_index(sps_j, pcoords))->
projs<2>();
286 post_elems[j] = p_val;
289 auto [post_mem, elem_post] = after_post->
vars<2>();
290 auto stored =
buffer::op_write(obr, obs, obT, post_mem, wu, fold_index(So, write_coords), elem_post);
291 after_post->app(
true, wk,
w.tuple({stored->proj(0), stored->proj(1)}));
292 apply_cps(wb, post, {pcur, wv,
w.tuple(post_elems)}, after_post);
297 fun->app(
true,
w.app(nest,
w.tuple({init, cell, wb})),
w.tuple({a_mem, out_buf, cont}));
302const Def* LowerMapReduce::lower_broadcast(
const App* app) {
304 auto callee = app->callee()->as<
App>();
305 auto [s_in, s_out] =
rewrite(callee->arg())->
projs<2>();
307 auto result_ty =
rewrite(app->type());
309 auto r_nat = s_out->num_projs();
311 auto mem_ty =
w.call<
mem::M>(0);
312 auto fun =
w.mut_fun(
w.sigma({mem_ty, input->type()}), result_ty)->set(
"broadcast");
314 auto [fun_mem, in_buf] = fun->var(0_n)->projs<2>();
315 auto cont = fun->var(1);
321 const Def* acc =
w.tuple({a_mem, out_buf});
322 auto current_mut = fun;
324 out_iters.reserve(r_nat);
325 for (
size_t i = 0; i < r_nat; ++i) {
326 auto dim = s_out->proj(r_nat, i);
328 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(
"bcast_" + std::to_string(i)));
329 auto [iter, new_acc, yield] = body->vars<3>();
333 current_mut->set(
true, for_call);
336 auto [loop_mem, loop_buf] = acc->projs<2>();
339 auto iters =
w.tuple(out_iters);
340 auto [rd_mem, rd_val] =
buffer::op_read(in_r, in_s, in_T, loop_mem, in_buf, fold_index(s_in, iters))->
projs<2>();
341 auto [wr_mem, wr_buf]
342 =
buffer::op_write(out_r, out_s, out_T, rd_mem, loop_buf, fold_index(s_out, iters), rd_val)->
projs<2>();
343 current_mut->app(
true, cont,
w.tuple({wr_mem, loop_buf}));
348const Def* LowerMapReduce::lower_pad(
const App* app) {
354 auto [Tr, s_in, params, s_out] =
c->uncurry_args<4>();
355 auto [
mode, lo, hi] = params->projs<3>();
356 auto [op_mem, input, value] =
rewrite(app->arg())->
projs<3>();
357 auto result_ty =
rewrite(app->type());
361 if (!r_l || !mode_l) {
362 log().w(
"rank/mode of {} is not known at lowering time", app);
363 return RWPhase::rewrite_imm_App(app);
366 auto mode_nat = *mode_l;
367 auto i64 =
w.type_i64();
370 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
372 auto compute = [&](
const DefVec& iters,
const Def* ins,
const Def* mem) -> std::pair<const Def*, const Def*> {
373 auto [in_buf, fill] = ins->projs<2>();
377 for (
u64 d = 0;
d < rn; ++
d) {
384 valid.push_back(v_d);
385 idx_i64 = sel(v_d, in_d,
w.lit_i64(0));
389 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({in_d, sin_m1}))}));
395 if (mode_nat != 0)
return {rd_mem, elem};
396 auto all_valid = valid.empty() ?
w.lit_tt() : valid[0];
397 for (
u64 d = 1;
d < valid.size(); ++
d)
399 return {rd_mem, sel(all_valid, elem, fill)};
402 return build_pointwise(w, result_ty, op_mem,
w.tuple({input, value}), s_out, rn,
"pad", compute);
405const Def* LowerMapReduce::lower_concat(
const App* app) {
406 auto&
w = new_world();
407 auto c = rewrite(app->callee())->as<
App>();
411 auto [TnisR, ax, Sis, s_out] =
c->uncurry_args<4>();
412 auto [T, nis,
r] = TnisR->projs<3>();
413 auto [op_mem, op_is] = rewrite(app->arg())->projs<2>();
414 auto result_ty = rewrite(app->type());
419 if (!nis_l || !r_l || !ax_l) {
420 log().w(
"nis/rank/axis of {} are not known at lowering time", app);
421 return RWPhase::rewrite_imm_App(app);
423 auto nisn = *nis_l, rn = *r_l, axn = *ax_l;
427 fe::Vector<u64> ext(nisn);
429 for (
u64 i = 0; i < nisn; ++i) {
430 off[i] =
w.lit_i64(acc_off);
433 log().w(
"extent of input {} of {} along the concat axis is not known at lowering time", i, app);
434 return RWPhase::rewrite_imm_App(app);
440 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
442 auto compute = [&](
const DefVec& iters,
const Def* ins,
const Def* mem) -> std::pair<const Def*, const Def*> {
443 auto o_ax = iters[axn];
444 const Def* cur = mem;
446 auto read_i = [&](
u64 i) ->
const Def* {
447 auto in_buf = ins->proj(nisn, i);
449 auto Sis_i = Sis->proj(nisn, i);
450 auto e_i_m1 =
w.lit_i64(ext[i] - 1);
453 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({loc, e_i_m1}))}));
455 for (
u64 d = 0;
d < rn; ++
d) {
456 auto idx_i64 = (
d == axn) ? clamp : iters[
d];
459 auto [rd_mem, rd_val]
465 auto result = read_i(0);
466 for (
u64 i = 1; i < nisn; ++i) {
468 result = sel(cond, read_i(i), result);
470 return {cur, result};
473 return build_pointwise(w, result_ty, op_mem, op_is, s_out, rn,
"concat", compute);
476const Def* LowerMapReduce::lower_gather(
const App* app) {
477 auto&
w = new_world();
478 auto c = rewrite(app->callee())->as<
App>();
480 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
481 auto [T,
r] = Tr->projs<2>();
482 auto [s_src, s_idx] = shapes->projs<2>();
483 auto [op_mem, input,
idx] = rewrite(app->arg())->projs<3>();
484 auto result_ty = rewrite(app->type());
488 if (!r_l || !dim_l) {
489 log().w(
"{} doesn't have lowering-time known rank/axis", app);
490 return RWPhase::rewrite_imm_App(app);
492 auto rn = *r_l, axis = *dim_l;
494 auto compute = [&](
Defs iters,
const Def* ins,
const Def* mem) -> std::pair<const Def*, const Def*> {
495 auto [in_buf, index_buf] = ins->projs<2>();
500 for (
u64 d = 0;
d < rn; ++
d)
501 index_coords[d] =
w.call(
core::conv::u, s_idx->proj(rn, d), iters[d]);
502 auto [index_mem, selected]
503 =
buffer::op_read(xbr, xbs, xbT, mem, index_buf, fold_index(s_idx,
w.tuple(index_coords)))->
projs<2>();
507 for (
u64 d = 0;
d < rn; ++
d) {
508 auto coordinate =
d == axis ? selected_i64 : iters[
d];
509 source_coords[
d] =
w.call(
core::conv::u, s_src->proj(rn, d), coordinate);
511 auto [read_mem, value]
512 =
buffer::op_read(ibr, ibs, ibT, index_mem, in_buf, fold_index(s_src,
w.tuple(source_coords)))->
projs<2>();
513 return {read_mem, value};
515 return build_pointwise(w, result_ty, op_mem,
w.tuple({input, idx}), s_idx, rn,
"gather", compute);
518const Def* LowerMapReduce::lower_scatter(
const App* app) {
519 auto&
w = new_world();
520 auto c = rewrite(app->callee())->as<
App>();
522 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
523 auto [T,
r] = Tr->projs<2>();
524 auto [s_src, s_idx, s_updates] = shapes->projs<3>();
525 auto [op_mem, input,
idx, updates] = rewrite(app->arg())->projs<4>();
526 auto result_ty = rewrite(app->type());
530 if (!r_l || !dim_l) {
531 log().w(
"{} doesn't have lowering-time known rank/axis", app);
532 return RWPhase::rewrite_imm_App(app);
534 auto rn = *r_l, axis = *dim_l;
536 auto mem_ty =
w.call<
mem::M>(0);
537 auto fun =
w.mut_fun(
w.sigma({mem_ty, input->type(), idx->type(), updates->type()}), result_ty)->set(
"scatter");
539 auto [fun_mem, in_buf, index_buf, update_buf] = fun->var(0_n)->projs<4>();
540 auto cont = fun->var(1);
544 auto copy_mem =
buffer::op_copy(obr, obs, obT, a_mem, out_buf, in_buf);
545 const Def* acc =
w.tuple({copy_mem, out_buf});
549 for (
u64 d = 0;
d < rn; ++
d) {
551 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(
"scatter_" + std::to_string(d)));
552 auto [iter, new_acc, yield] = body->vars<3>();
555 iters.push_back(iter);
556 current->set(
true, for_call);
559 auto [loop_mem, loop_buf] = acc->projs<2>();
564 for (
u64 d = 0;
d < rn; ++
d)
565 index_coords[d] =
w.call(
core::conv::u, s_idx->proj(rn, d), iters[d]);
566 auto folded_index = fold_index(s_idx,
w.tuple(index_coords));
568 for (
u64 d = 0;
d < rn; ++
d)
569 update_coords[d] =
w.call(
core::conv::u, s_updates->proj(rn, d), iters[d]);
570 auto [index_mem, selected] =
buffer::op_read(xbr, xbs, xbT, loop_mem, index_buf, folded_index)->
projs<2>();
571 auto [update_mem, update]
572 =
buffer::op_read(ubr, ubs, ubT, index_mem, update_buf, fold_index(s_updates,
w.tuple(update_coords)))
576 DefVec destination_coords(rn);
577 for (
u64 d = 0;
d < rn; ++
d) {
578 auto coordinate =
d == axis ? selected_i64 : iters[
d];
579 destination_coords[
d] =
w.call(
core::conv::u, s_src->proj(rn, d), coordinate);
581 auto [write_mem, written]
582 =
buffer::op_write(obr, obs, obT, update_mem, loop_buf, fold_index(s_src,
w.tuple(destination_coords)), update)
584 current->app(
true, cont,
w.tuple({write_mem, loop_buf}));
const Def * callee() const
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...
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
Lam * set(Filter filter, const Def *body)
static std::optional< T > isa(const Def *def)
const fe::Log & log() const
const fe::Vector< std::string > & args()
Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>.
bool is_bootstrapping() const
Returns whether we are currently bootstrapping (rewriting annexes).
World & new_world()
Create new Defs into this.
virtual const Def * rewrite(const Def *)
const Def * rewrite_imm_App(const App *) override
const Def * op_write(const Def *r, const Def *s, const Def *T, const Def *mem, const Def *buf, const Def *idx, const Def *val)
buffer.write (r, s, T) (mem, buf, idx, val) ↦ [mem.M 0, buffer.Buf (r, s, T)].
const Def * op_read(const Def *r, const Def *s, const Def *T, const Def *mem, const Def *buf, const Def *idx)
buffer.read (r, s, T) (mem, buf, idx) ↦ [mem.M 0, T].
const Def * op_alloc(const Def *r, const Def *s, const Def *T, const Def *mem)
buffer.alloc (r, s, T) mem ↦ [mem.M 0, buffer.Buf (r, s, T)].
const Def * op_copy(const Def *r, const Def *s, const Def *T, const Def *mem, const Def *dst, const Def *src)
buffer.copy (r, s, T) (mem, dst, src) ↦ mem.M 0 (copies the whole buffer src into dst).
const Def * op_cps2ds_dep(const Def *k)
Lam * mut_con(World &w, nat_t a=0)
Yields con[mem.M 0].
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec