20const Def* LowerMapReduce::rec_broadcast(
const Def* s_in,
const Def* s_out,
const Def* input,
u64 r,
u64 i) {
23 if (i == r)
return input;
25 auto s_in_ri = s_in->proj(r, i), s_out_ri = s_out->proj(r, i);
26 log().d(
"broadcast dimension {} of {}: {} → {}, input = {}: {}", i, r, s_in_ri, s_out_ri, input, input->type());
28 if (s_in_ri == s_out_ri) {
31 [&](
size_t j) {
return rec_broadcast(s_in, s_out, input->proj(*s_in_lit, j), r, i + 1); });
32 return w.tuple(inputs);
36 log().w(
"dimension {} has equal but non-literal extent: {}", i, s_in_ri);
41 if (
auto s_in_lit =
Lit::isa<u64>(s_in_ri); s_in_lit && *s_in_lit == 1) {
42 log().d(
"dimension {}: packing the size-1 input to {}", i, s_out_ri);
43 return w.pack(s_out_ri, rec_broadcast(s_in, s_out, input, r, i + 1));
46 log().w(
"cannot broadcast dimension {}: {} → {}", i, s_in_ri, s_out_ri);
50const Def* LowerMapReduce::lower_broadcast(
const App* app) {
55 auto [s_in, s_out, input] = arg->projs<3>();
56 auto callee =
c->as<
App>();
57 auto [T,
r] = callee->args<2>();
58 log().d(
"lower broadcast: input = {}: {}, T = {}, r = {}, s_in = {}, s_out = {}", input, input->type(), T, r, s_in,
63 log().w(
"rank {} of {} is not known at lowering time", r, app);
67 if (s_in == s_out)
return input;
71 assert(*s_in_lit == 1 &&
"input dimensions must be 1 or equal to the output dimension");
72 return w.pack(s_out, input);
76 auto result = rec_broadcast(s_in, s_out, input, *r_nat, 0);
77 log().d(
"broadcast result: {}", result);
83std::pair<Lam*, const Def*> counting_for(
const Def* bound,
const Def* acc,
const Def* exit, Sym name) {
84 auto&
w = bound->world();
85 auto acc_ty = acc->type();
86 auto body =
w.mut_con({
w.type_i64(), acc_ty,
w.cn(acc_ty)})->
set(name);
87 auto for_loop =
w.call<
affine::For>(body, exit,
Defs{
w.lit_i64(0), bound,
w.lit_i64(1), acc});
88 return {body, for_loop};
93DefVec build_loops(World& w, Lam*& cur,
const Def*& exit,
const Def*& acc,
Defs dims, std::string_view name) {
95 iters.reserve(dims.size());
96 for (
size_t i = 0, e = dims.size(); i != e; ++i) {
98 auto [body, for_call] = counting_for(bound, acc, exit,
w.sym(std::format(
"{}_{}", name, i)));
99 auto [iter, new_acc, yield] = body->vars<3>();
102 iters.emplace_back(iter);
103 cur->set(
true, for_call);
109const Def* elem_type(
const Def* type,
u64 r) {
110 for (
u64 i = 0; i !=
r; ++i)
111 if (
auto seq =
type->isa<Seq>())
118const Def* nested_extract(World& w,
const Def* matrix,
const Def* coords,
const Def*
shape,
u64 r) {
119 return op_get(elem_type(matrix->type(), r),
w.lit_nat(r),
shape, matrix, coords);
122const Def* nested_insert(World& w,
const Def* matrix,
const Def* coords,
const Def*
shape,
u64 r,
const Def* elem) {
123 return op_set(elem_type(matrix->type(), r),
w.lit_nat(r),
shape, matrix, coords, elem);
127std::optional<fe::Vector<u64>> lit_projs(
const Def* def,
u64 n) {
128 auto res = fe::Vector<u64>(n);
129 for (
u64 i = 0; i != n; ++i)
138const Def*
select(World& w,
const Def* cond,
const Def* t,
const Def* f) {
return w.extract(
w.tuple({f, t}), cond); }
141const Def* clamp(World& w,
const Def* x,
const Def* bound) {
143 return w.call(
core::extrema::smax,
w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({x, hi}))}));
148const Def* LowerMapReduce::lower_map_reduce(
const App* app) {
165 auto&
w = new_world();
166 auto c = rewrite(app->callee())->as<
App>();
167 auto inputs = rewrite(app->arg());
168 auto type = rewrite(app->type());
170 auto [nis_nps, meta, shapes, in_tys, comb_init, acc_out, accs_all] =
c->uncurry_args<7>();
171 auto [nis, nps] = nis_nps->projs<2>();
172 auto [To, Tp, Ro, Rn, TSched] = meta->projs<5>();
173 auto [So, Sr, sched] = shapes->projs<3>();
174 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->projs<6>();
175 auto [comb,
init, post] = comb_init->projs<3>();
176 auto [accs, post_accs] = accs_all->projs<2>();
181 if (!nis_l || !nps_l || !ro_l || !rn_l || *rn_l < *ro_l) {
182 log().w(
"rank counts (nis/nps/Ro/Rn) of {} are not known at lowering time", app);
185 auto nis_nat = *nis_l;
186 auto nps_nat = *nps_l;
187 auto ro = *ro_l, rr = *rn_l - *ro_l;
189 auto n =
w.lit_nat(nloops);
192 auto ris_nat = lit_projs(Ris, nis_nat);
193 auto rps_nat = lit_projs(Rps, nps_nat);
194 if (!ris_nat || !rps_nat) {
195 log().w(
"the input ranks of {} are not known at lowering time", app);
202 auto mem0 =
w.app(
w.annex<
mem::M>(),
w.lit_nat(0));
203 auto affine_map = [&](
const Def*
f,
const Def* m,
const Def* n,
const Def*
sin,
const Def* sout,
const Def* idxs) {
205 a =
w.app(a,
w.tuple({sin, sout}));
208 a =
w.app(a,
w.lit_nat_0());
209 return w.app(a,
w.bot(mem0))->proj(2, 1);
213 auto fun =
w.mut_fun(inputs->type(), type)->set(
"mapRed");
215 auto call =
w.app(ds_fun, inputs)->set(
"call");
217 auto [new_inputs, cont] = fun->vars<2>();
218 auto [new_is, new_post_is] = new_inputs->set(
"is")->projs<2>();
219 auto sr = Sr->projs(nloops);
222 const Def* acc =
w.bot(cont->type()->as<Pi>()->dom());
223 auto current_mut = fun;
224 auto raw_out = build_loops(w, current_mut, cont, acc,
Defs(sr).subspan(0, ro),
"forOut");
226 auto wb_matrix = acc;
232 auto write_back =
w.mut_con(To)->set(
"writeBack");
233 auto element_final = write_back->var();
234 DefVec wb_iters = out_iters;
235 for (
u64 j = 0; j < rr; ++j)
236 wb_iters.emplace_back(
w.call(
core::conv::u, sr[ro + j],
w.lit_i64(0)));
237 auto write_coords = affine_map(acc_out, Ro, n, Sr, So,
w.tuple(wb_iters));
240 DefVec post_elements(nps_nat, [&](
size_t j) {
241 auto sps_j = Sps->proj(nps_nat, j);
242 auto coords = affine_map(post_accs->proj(nps_nat, j), Rps->proj(nps_nat, j), Ro, So, sps_j, write_coords);
243 return nested_extract(w, new_post_is->proj(nps_nat, j), coords, sps_j, (*rps_nat)[j]);
246 auto after_post =
w.mut_con(Tp)->set(
"afterPost");
247 after_post->app(
true, cont, nested_insert(w, wb_matrix, write_coords, So, ro, after_post->var()));
248 write_back->app(
true, post, {
w.tuple({element_final,
w.tuple(post_elements)}), after_post});
254 auto raw_red = build_loops(w, current_mut, cont, acc,
Defs(sr).subspan(ro, rr),
"forIn");
255 auto element_acc = acc;
258 DefVec iters_v = out_iters;
259 for (
u64 j = 0; j != rr; ++j)
260 iters_v.emplace_back(
w.call(
core::conv::u, sr[ro + j], raw_red[j]));
261 auto iters =
w.tuple(iters_v);
264 DefVec input_elements(nis_nat, [&](
size_t i) {
265 auto sis_i = Sis->proj(nis_nat, i);
266 auto coords = affine_map(accs->proj(nis_nat, i), Ris->proj(nis_nat, i), n, Sr, sis_i, iters);
267 return nested_extract(w, new_is->proj(nis_nat, i), coords, sis_i, (*ris_nat)[i]);
272 current_mut->app(
true, comb, {
w.tuple({element_acc,
w.tuple(input_elements)}), cont});
274 }
catch (
const std::exception& e) { fe::throwf(
"failed to lower `tensor.map_reduce`: {}",
e.what()); }
277const Def* LowerMapReduce::build_pointwise(
const Def* inputs,
281 std::function<
const Def*(
Defs,
const Def*)> compute) {
282 auto&
w = new_world();
284 auto fun =
w.mut_fun(inputs->type(), type)->set(
"pointwise");
286 auto call =
w.app(ds_fun, inputs)->set(
"call");
288 auto [new_inputs, cont] = fun->vars<2>();
289 new_inputs->set(
"is");
292 const Def* acc =
w.bot(cont->type()->as<Pi>()->dom());
293 auto current_mut = fun;
294 auto so = So->projs(ro);
295 auto out_iters = build_loops(w, current_mut, cont, acc, so,
"forOut");
296 auto wb_matrix = acc;
300 auto element = compute(out_iters, new_inputs);
301 current_mut->app(
true, cont, nested_insert(w, wb_matrix,
w.tuple(write_coords), So, ro, element));
305const Def* LowerMapReduce::lower_generate(
const App* app) {
306 auto&
w = new_world();
307 auto c = rewrite(app->callee())->as<
App>();
308 auto body = rewrite(app->arg());
309 auto type = rewrite(app->type());
311 auto [meta, s_out] =
c->uncurry_args<2>();
312 auto [T,
r] = meta->projs<2>();
315 log().w(
"rank {} of {} is not known at lowering time", r, app);
322 if (!
type->isa<Arr>()) {
323 DefVec zeros(rn, [&](
size_t) {
return w.lit_i64(0); });
324 return w.app(body,
w.tuple(zeros));
327 auto unit =
w.tuple(
Defs{});
328 auto compute = [&](
Defs out_iters,
const Def*) {
return w.call(body, out_iters); };
329 return build_pointwise(unit, type, s_out, rn, compute);
332const Def* LowerMapReduce::lower_pad(
const App* app) {
333 auto&
w = new_world();
334 auto c = rewrite(app->callee())->as<
App>();
335 auto args = rewrite(app->arg());
336 auto type = rewrite(app->type());
339 auto [Tr, s_in, params] =
c->uncurry_args<3>();
340 auto [T,
r] = Tr->projs<2>();
341 auto [
mode, lo, hi] = params->projs<3>();
345 if (!r_l || !mode_l) {
346 log().w(
"rank/mode of {} is not known at lowering time", app);
350 auto mode_nat = *mode_l;
351 auto i64 =
w.type_i64();
355 auto inner_type =
type;
356 for (
u64 d = 0;
d < rn; ++
d) {
357 auto inner_type_seq = inner_type->as<Seq>();
358 so[
d] = inner_type_seq->arity();
359 inner_type = inner_type_seq->body();
361 auto s_out =
w.tuple(so);
363 auto compute = [&](
Defs out_iters,
const Def* new_inputs) ->
const Def* {
364 auto [input, value] = new_inputs->projs<2>();
367 for (
u64 d = 0;
d < rn; ++
d) {
374 valid.push_back(v_d);
375 idx_i64 =
select(w, v_d, in_d,
w.lit_i64(0));
377 idx_i64 = clamp(w, in_d, sin_d);
381 auto elem = nested_extract(w, input,
w.tuple(clamped), s_in, rn);
382 if (mode_nat != 0)
return elem;
383 auto all_valid = valid.empty() ?
w.lit_tt() : valid[0];
384 for (
u64 d = 1;
d < valid.size(); ++
d)
386 return select(w, all_valid, elem, value);
389 return build_pointwise(args, type, s_out, rn, compute);
392const Def* LowerMapReduce::lower_concat(
const App* app) {
393 auto&
w = new_world();
394 auto c = rewrite(app->callee())->as<
App>();
395 auto args = rewrite(app->arg());
396 auto type = rewrite(app->type());
399 auto [TnisR, ax, Sis] =
c->uncurry_args<3>();
400 auto [T, nis,
r] = TnisR->projs<3>();
405 if (!nis_l || !r_l || !ax_l) {
406 log().w(
"nis/r/ax of {} are not known at lowering time", app);
409 auto nisn = *nis_l, rn = *r_l, axn = *ax_l;
410 auto i64 =
w.type_i64();
415 for (
u64 i = 0; i < nisn; ++i) {
416 off[i] =
w.lit_i64(acc_off);
419 log().w(
"extent of input {} of {} along the concat axis is not known at lowering time", i, app);
427 for (
u64 d = 0;
d < rn; ++
d)
428 so[d] = (d == axn) ?
w.lit_nat(acc_off) : Sis->proj(nisn, 0)->proj(rn, d);
429 auto s_out =
w.tuple(so);
431 auto compute = [&](
Defs out_iters,
const Def* new_inputs) ->
const Def* {
432 auto o_ax = out_iters[axn];
434 auto read_i = [&](
u64 i) ->
const Def* {
435 auto Sis_i = Sis->proj(nisn, i);
438 auto in_ax = clamp(w, loc, e_i);
439 DefVec coords(rn, [&](
size_t d) {
440 return w.call(
core::conv::u, Sis_i->proj(rn, d), d == axn ? in_ax : out_iters[d]);
442 return nested_extract(w, new_inputs->proj(nisn, i),
w.tuple(coords), Sis_i, rn);
445 auto result = read_i(0);
446 for (
u64 i = 1; i < nisn; ++i) {
448 result =
select(w, cond, read_i(i), result);
453 return build_pointwise(args, type, s_out, rn, compute);
456const Def* LowerMapReduce::lower_gather(
const App* app) {
457 auto&
w = new_world();
458 auto c = rewrite(app->callee())->as<
App>();
459 auto args = rewrite(app->arg());
460 auto type = rewrite(app->type());
462 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
463 auto [T,
r] = Tr->projs<2>();
464 auto [s_src, s_idx] = shapes->projs<2>();
467 if (!r_l || !dim_l)
return nullptr;
474 auto compute = [&](
Defs out_indices,
const Def* inputs) ->
const Def* {
476 return w.call(element,
Defs{s_src, s_idx}, dim, out_indices, inputs);
478 return build_pointwise(args, type, s_idx, *r_l, compute);
481const Def* LowerMapReduce::lower_scatter(
const App* app) {
482 auto&
w = new_world();
483 auto c = rewrite(app->callee())->as<
App>();
484 auto args = rewrite(app->arg());
485 auto type = rewrite(app->type());
487 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
488 auto [T,
r] = Tr->projs<2>();
489 auto [s_src, s_idx, s_updates] = shapes->projs<3>();
492 if (!r_l || !dim_l)
return nullptr;
499 auto fun =
w.mut_fun(args->type(), type)->set(
"scatter");
501 auto call =
w.app(ds_fun, args)->set(
"call");
503 auto [iiu, cont] = fun->vars<2>();
504 auto [input, index, updates] = iiu->projs<3>();
507 auto dims = s_idx->projs(rn);
508 auto visit_indices = build_loops(w, current, cont, acc, dims,
"scatter");
511 auto next =
w.call(step,
Defs{s_src, s_idx, s_updates}, dim, visit_indices,
Defs{acc, index, updates});
512 current->app(
true, cont, next);
521 if (
auto res = lower_broadcast(bc))
return res;
523 if (
auto res = lower_map_reduce(mr))
return res;
525 if (
auto res = lower_generate(
generate))
return res;
527 if (
auto res = lower_pad(
pad))
return res;
529 if (
auto res = lower_concat(
cat))
return res;
531 if (
auto res = lower_gather(
gather))
return res;
533 if (
auto res = lower_scatter(
scatter))
return res;
535 return RWPhase::rewrite_imm_App(app);
static auto isa(const Def *def)
Def * set(size_t i, const Def *)
Successively set from left to right.
static std::optional< T > isa(const Def *def)
const fe::Log & log() const
World & new_world()
Create new Defs into this.
virtual const Def * rewrite(const Def *)
const Def * rewrite_imm_App(const App *) final
const Def * op_cps2ds_dep(const Def *k)
bool check_scatter_shape_constraints(const Def *rank, const Def *dim, const Def *source_shape, const Def *index_shape, const Def *updates_shape)
Checks statically decidable scatter constraints. Returns false for unresolved relations.
bool check_gather_shape_constraints(const Def *rank, const Def *dim, const Def *source_shape, const Def *index_shape)
Checks statically decidable gather constraints. Returns false for unresolved relations.
const Def * op_set(const Def *T, const Def *r, const Def *s, const Def *arr, const Def *index, const Def *x)
gather_pointwise_elem_impl
const Def * op_get(const Def *T, const Def *r, const Def *s, const Def *arr, const Def *index)
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec