16const Def* LowerMapReduce::rec_broadcast(
const Def* s_in,
const Def* s_out,
const Def* input,
u64 r,
u64 i) {
19 if (i == r)
return input;
21 auto s_in_ri = s_in->proj(r, i), s_out_ri = s_out->proj(r, i);
22 DLOG(
"rec_broadcast");
25 DLOG(
" s_in_ri = {} : {}", s_in_ri, s_in_ri->type());
26 DLOG(
" s_out_ri = {} : {}", s_out_ri, s_out_ri->type());
27 DLOG(
" input = {} : {}", input, input->type());
29 if (s_in_ri == s_out_ri) {
31 DefVec inputs(*s_in_lit, [&](
size_t j) {
return rec_broadcast(s_in, s_out, input->proj(j), r, i + 1); });
32 return w.tuple(inputs);
36 WLOG(
"dimension {} of the input and output are equal but not literal: {} : {}", i, s_in_ri,
42 if (
auto s_in_lit =
Lit::isa<u64>(s_in_ri); s_in_lit && *s_in_lit == 1) {
43 DLOG(
"dimension {} of the input is 1, can be broadcasted to dimension {} of the output", i, s_out_ri);
44 return w.pack(s_out_ri, rec_broadcast(s_in, s_out, input, r, i + 1));
47 WLOG(
"cannot broadcast dimension {} of size {} to size {}", i, s_in_ri, s_out_ri);
51const Def* LowerMapReduce::lower_broadcast(
const App* app) {
56 auto [s_in, s_out, input] = arg->projs<3>();
57 auto callee =
c->as<
App>();
58 auto [T,
r] = callee->args<2>();
59 DLOG(
"lower_broadcast");
60 DLOG(
" s_out = {} : {}", s_out, s_out->type());
61 DLOG(
" input = {} : {}", input, input->type());
62 DLOG(
" T = {} : {}", T, T->type());
63 DLOG(
" r = {} : {}", r,
r->type());
64 DLOG(
" s_in = {} : {}", s_in, s_in->type());
68 WLOG(
"{} doesn't have a lowering-time known rank: {}", app, r);
72 if (s_in == s_out)
return input;
76 assert(*s_in_lit == 1 &&
"input dimensions must be 1 or equal to the output dimension");
77 return w.pack(s_out, input);
81 auto result = rec_broadcast(s_in, s_out, input, *r_nat, 0);
82 DLOG(
"result of rec_broadcast = {} : {}", result, result->type());
87 auto& w = bound->
world();
88 auto acc_ty = acc->
type();
89 auto body = w.mut_con({ w.type_i64(), acc_ty, w.cn(acc_ty)})->
set(name);
90 auto for_loop = w.call<
affine::For>(body, exit,
Defs{w.lit_i64(0), bound, w.lit_i64(1), acc});
91 return {body, for_loop};
96 for (
u64 i = 0; i < r; ++i)
97 if (
auto seq = cur->isa<
Seq>())
115const Def* LowerMapReduce::lower_map_reduce(
const App* app) {
134 auto [nis, meta, shapes, TisRisSis, comb_init, acc_out, accs] = c->uncurry_args<7>();
135 auto [To, Ro, Rr] = meta->projs<3>();
136 auto [So, Sr] = shapes->projs<2>();
137 auto [Tis, Ris, Sis] = TisRisSis->projs<3>();
138 auto [comb, init] = comb_init->projs<2>();
142 if (!nis_l || !ro_l || !rr_l) {
143 WLOG(
"{} doesn't have lowering-time known rank counts (nis/Ro/Rr)", app);
146 auto nis_nat = *nis_l;
147 auto ro = *ro_l, rr = *rr_l;
148 auto nloops = ro + rr;
149 auto n = w.lit_nat(nloops);
153 for (
u64 i = 0; i < nis_nat; ++i) {
156 WLOG(
"input {} of {} has a non-literal rank", i, app);
165 auto mem0 = w.app(w.annex<
mem::M>(), w.lit_nat(0));
166 auto affine_map = [&](
const Def* f,
const Def* m,
const Def* n,
const Def* sin,
const Def* sout,
const Def* idxs) {
167 auto a = w.app(w.annex<
affine::map>(), w.tuple({m, n}));
168 a = w.app(a, w.tuple({sin, sout}));
171 return w.app(a, w.bot(mem0))->proj(2, 1);
175 auto fun =
w.mut_fun(inputs->
type(), type)->set(
"mapRed");
177 auto call =
w.app(ds_fun, inputs)->set(
"call");
179 auto new_inputs = fun->var(0)->set(
"is");
182 auto cont = fun->var(1);
183 auto init_mat =
w.bot(cont->type()->as<
Pi>()->dom());
185 auto current_mut = fun;
187 out_iters.reserve(ro);
188 for (
u64 i = 0; i < ro; ++i) {
189 auto dim = Sr->proj(nloops, i);
191 auto [body, for_call] =
counting_for(bound, acc, cont,
w.sym(
"forOut_" + std::to_string(i)));
192 auto [iter, new_acc, yield] = body->vars<3>();
196 current_mut->set(
true, for_call);
199 auto wb_matrix = acc;
204 auto write_back =
w.mut_con(To)->set(
"writeBack");
205 auto element_final = write_back->var(0);
206 DefVec wb_iters = out_iters;
207 for (
u64 j = 0; j < rr; ++j)
208 wb_iters.push_back(
w.call(
core::conv::u, Sr->proj(nloops, ro + j),
w.lit(
w.type_i64(), 0)));
209 auto write_coords = affine_map(acc_out, Ro, n, Sr, So,
w.tuple(wb_iters));
210 write_back->app(
true, cont,
nested_insert(w, wb_matrix, write_coords, So, ro, element_final));
216 red_iters.reserve(rr);
217 for (
u64 j = 0; j < rr; ++j) {
218 auto dim = Sr->proj(nloops, ro + j);
220 auto [body, for_call] =
counting_for(bound, acc, cont,
w.sym(
"forIn_" + std::to_string(j)));
221 auto [iter, new_acc, yield] = body->vars<3>();
225 current_mut->set(
true, for_call);
228 auto element_acc = acc;
231 DefVec iters_v = out_iters;
232 iters_v.insert(iters_v.end(), red_iters.begin(), red_iters.end());
233 auto iters =
w.tuple(iters_v);
236 DefVec input_elements(nis_nat);
237 for (
u64 i = 0; i < nis_nat; ++i) {
238 auto input_matrix = new_inputs->proj(nis_nat, i);
239 auto sis_i = Sis->proj(nis_nat, i);
240 auto coords = affine_map(accs->proj(nis_nat, i), Ris->proj(nis_nat, i), n, Sr, sis_i, iters);
241 input_elements[i] =
nested_extract(w, input_matrix, coords, sis_i, ris_nat[i]);
245 current_mut->app(
true, comb, {
w.tuple({element_acc,
w.tuple(input_elements)}), cont});
247 }
catch (
const std::exception& e) { fe::throwf(
"error during lowering map_reduce: {}",
e.what()); }
250const Def* LowerMapReduce::build_pointwise(
const Def* inputs,
254 std::function<
const Def*(
const DefVec&,
const Def*)> compute) {
257 auto fun =
w.mut_fun(inputs->type(), type)->set(
"pointwise");
259 auto call =
w.app(ds_fun, inputs)->set(
"call");
261 auto new_inputs = fun->var(0)->set(
"is");
264 auto cont = fun->var(1);
265 auto acc =
w.bot(cont->type()->as<
Pi>()->dom());
266 auto current_mut = fun;
268 out_iters.reserve(ro);
269 for (
u64 i = 0; i < ro; ++i) {
270 auto dim = So->proj(ro, i);
272 auto [body, for_call] =
counting_for(bound, acc, cont,
w.sym(
"forOut_" + std::to_string(i)));
273 auto [iter, new_acc, yield] = body->vars<3>();
275 out_iters.push_back(iter);
277 current_mut->set(
true, for_call);
280 auto wb_matrix = acc;
284 for (
u64 i = 0; i < ro; ++i)
285 write_coords[i] =
w.call(
core::conv::u, So->proj(ro, i), out_iters[i]);
286 auto element = compute(out_iters, new_inputs);
287 current_mut->app(
true, cont,
nested_insert(w, wb_matrix,
w.tuple(write_coords), So, ro, element));
291const Def* LowerMapReduce::lower_pad(
const App* app) {
298 auto [Tr, s_in, params] =
c->uncurry_args<3>();
299 auto [T,
r] = Tr->projs<2>();
300 auto [
mode, lo, hi] = params->projs<3>();
304 if (!r_l || !mode_l) {
305 WLOG(
"{} doesn't have a lowering-time known rank/mode", app);
309 auto mode_nat = *mode_l;
310 auto i64 =
w.type_i64();
314 auto inner_type =
type;
315 for (
u64 d = 0;
d < rn; ++
d) {
316 auto inner_type_seq = inner_type->as<Seq>();
317 so[
d] = inner_type_seq->arity();
318 inner_type = inner_type_seq->body();
320 auto s_out =
w.tuple(so);
323 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
325 auto compute = [&](
const DefVec& out_iters,
const Def* new_inputs) ->
const Def* {
326 auto [input, value] = new_inputs->projs<2>();
329 for (
u64 d = 0;
d < rn; ++
d) {
336 valid.push_back(v_d);
337 idx_i64 = sel(v_d, in_d,
w.lit_i64(0));
341 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({in_d, sin_m1}))}));
346 if (mode_nat != 0)
return elem;
347 auto all_valid = valid.empty() ?
w.lit_tt() : valid[0];
348 for (
u64 d = 1;
d < valid.size(); ++
d)
350 return sel(all_valid, elem, value);
353 return build_pointwise(args, type, s_out, rn, compute);
356const Def* LowerMapReduce::lower_concat(
const App* app) {
357 auto&
w = new_world();
358 auto c = rewrite(app->callee())->as<
App>();
359 auto args = rewrite(app->arg());
360 auto type = rewrite(app->type());
363 auto [TnisR, ax, Sis] =
c->uncurry_args<3>();
364 auto [T, nis,
r] = TnisR->projs<3>();
369 if (!nis_l || !r_l || !ax_l) {
370 WLOG(
"{} doesn't have lowering-time known nis/r/ax", app);
373 auto nisn = *nis_l, rn = *r_l, axn = *ax_l;
374 auto i64 =
w.type_i64();
379 for (
u64 i = 0; i < nisn; ++i) {
380 off[i] =
w.lit_i64(acc_off);
383 WLOG(
"{} input {} has a non-literal extent along the concat axis", app, i);
391 for (
u64 d = 0;
d < rn; ++
d)
392 so[d] = (d == axn) ?
w.lit_nat(acc_off) : Sis->proj(nisn, 0)->proj(rn, d);
393 auto s_out =
w.tuple(so);
395 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
397 auto compute = [&](
const DefVec& out_iters,
const Def* new_inputs) ->
const Def* {
398 auto o_ax = out_iters[axn];
400 auto read_i = [&](
u64 i) ->
const Def* {
401 auto Sis_i = Sis->proj(nisn, i);
406 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({loc, e_i_m1}))}));
408 for (
u64 d = 0;
d < rn; ++
d) {
409 auto idx_i64 = (
d == axn) ? clamp : out_iters[
d];
412 return nested_extract(w, new_inputs->proj(nisn, i),
w.tuple(coords), Sis_i, rn);
415 auto result = read_i(0);
416 for (
u64 i = 1; i < nisn; ++i) {
418 result = sel(cond, read_i(i), result);
423 return build_pointwise(args, type, s_out, rn, compute);
428 if (
auto res = lower_broadcast(bc))
return res;
432 if (
auto res = lower_pad(
pad))
return res;
434 if (
auto res = lower_concat(
cat))
return res;
436 return RWPhase::rewrite_imm_App(app);
const Def * callee() const
static auto isa(const Def *def)
Def * set(size_t i, const Def *)
Successively set from left to right.
World & world() const noexcept
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
static std::optional< T > isa(const Def *def)
const Vector< std::string > & args()
Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>.
World & new_world()
Create new Defs into this.
virtual const Def * rewrite(const Def *)
Base class for Arr and Pack.
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.
const Def * op_cps2ds_dep(const Def *k)
static const Def * nested_insert(World &w, const Def *matrix, const Def *coords, const Def *shape, u64 r, const Def *elem)
static const Def * nested_extract(World &w, const Def *matrix, const Def *coords, const Def *shape, u64 r)
static const Def * get_element_type(const Def *type, u64 r)
static std::pair< Lam *, const Def * > counting_for(const Def *bound, const Def *acc, const Def *exit, Sym name)
const Def * op_set(const Def *T, const Def *r, const Def *s, const Def *arr, const Def *index, const Def *x)
const Def * op_get(const Def *T, const Def *r, const Def *s, const Def *arr, const Def *index)
Vector< const Def * > DefVec