20const Def* fold_index(
const Def* shape,
const Def* idx) {
22 auto r =
shape->num_projs();
25 for (
size_t i = 0; i !=
r; ++i)
29 out.push_back(
idx->proj(r, i));
33 if (!dropped)
return idx;
38std::pair<Lam*, const Def*> counting_for(
const Def* bound,
const Def* acc,
const Def* exit, Sym name) {
39 auto&
w = bound->world();
40 auto acc_ty = acc->type();
41 auto body =
w.mut_con({
w.type_i64(), acc_ty,
w.cn(acc_ty)})->
set(name);
42 auto for_loop =
w.call<
affine::For>(body, exit,
Defs{
w.lit_i64(0), bound,
w.lit_i64(1), acc});
43 return {body, for_loop};
51template<
class Compute>
52const Def* build_pointwise(World& w,
58 const std::string& name,
60 auto mem_ty =
w.call<
mem::M>(0);
61 auto fun =
w.mut_fun(
w.sigma({mem_ty, op_ins->type()}), result_ty)->set(name);
63 auto [fun_mem, ins] = fun->var(0_n)->projs<2>();
64 auto cont = fun->var(1);
68 const Def* acc =
w.tuple({a_mem, out_buf});
69 auto current_mut = fun;
73 for (
u64 d = 0;
d < rn; ++
d) {
75 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(name +
"_" + std::to_string(d)));
76 auto [iter, new_acc, yield] = body->vars<3>();
78 iters.push_back(iter);
80 current_mut->set(
true, for_call);
83 auto [loop_mem, loop_buf] = acc->projs<2>();
85 std::pair<const Def*, const Def*> el = compute(iters, ins, loop_mem);
86 auto [el_mem, elem] = el;
89 for (
u64 d = 0;
d < rn; ++
d)
92 =
buffer::op_write(obr, obs, obT, el_mem, loop_buf, fold_index(s_out,
w.tuple(wcoords)), elem)->
projs<2>();
93 current_mut->app(
true, cont,
w.tuple({wr_mem, loop_buf}));
106 return RWPhase::rewrite_imm_App(app);
109const Def* LowerAff::lower_buffer_constant(
const App* app) {
116 if (!rn)
return RWPhase::rewrite_imm_App(app);
122 return build_pointwise(
123 w, result_ty,
mem, val, s_out, *rn,
"constant_fill",
124 [](
const DefVec&,
const Def* ins,
const Def* m) -> std::pair<const Def*, const Def*> {
return {m, ins}; });
127const Def* LowerAff::lower_map_reduce_aff(
const App* app) {
131 auto [nis, meta, shapes, TisRisSis, comb_init, acc_out, accs] = c->uncurry_args<7>();
132 auto [To, Ro, Rr] = meta->
projs<3>();
133 auto [So, Sr] = shapes->
projs<2>();
134 auto [Tis, Ris, Sis] = TisRisSis->
projs<3>();
135 auto [comb, init] = comb_init->
projs<2>();
143 if (!nis_l || !ro_l || !rr_l) {
144 WLOG(
"{} doesn't have lowering-time known rank counts (nis/Ro/Rr)", app);
145 return RWPhase::rewrite_imm_App(app);
147 auto nis_nat = *nis_l;
148 auto ro = *ro_l, rr = *rr_l;
149 auto nloops = ro + rr;
150 auto n = w.lit_nat(nloops);
154 auto affine_map = [&](
const Def* f,
const Def* m,
const Def* nn,
const Def* sin,
const Def* sout,
const Def* idxs,
156 auto a = w.app(w.annex<
affine::map>(), w.tuple({m, nn}));
157 a = w.app(a, w.tuple({sin, sout}));
160 return w.app(a, mem)->projs<2>();
163 auto mem_ty =
w.call<
mem::M>(0);
166 auto fun =
w.mut_fun(
w.sigma({mem_ty, op_is->type()}), result_ty)->set(
"mapRedAff");
168 auto [fun_mem, new_inputs] = fun->var(0_n)->projs<2>();
169 auto cont = fun->var(1);
174 const Def* acc =
w.tuple({a_mem, out_buf});
175 auto current_mut = fun;
178 out_iters.reserve(ro);
179 for (
u64 i = 0; i < ro; ++i) {
180 auto dim = Sr->proj(nloops, i);
182 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(
"forOut_" + std::to_string(i)));
183 auto [iter, new_acc, yield] = body->vars<3>();
187 current_mut->set(
true, for_call);
190 auto [wb_mem, wb_buf] = acc->projs<2>();
194 auto [wb_in_mem, elem_final] = write_back->
vars<2>();
195 DefVec wb_iters = out_iters;
196 for (
u64 j = 0; j < rr; ++j)
197 wb_iters.push_back(
w.call(
core::conv::u, Sr->proj(nloops, ro + j),
w.lit(
w.type_i64(), 0)));
198 auto [wc_mem, write_coords] = affine_map(acc_out, Ro, n, Sr, So,
w.tuple(wb_iters), wb_in_mem);
199 auto stored =
buffer::op_write(obr, obs, obT, wc_mem, wb_buf, fold_index(So, write_coords), elem_final);
200 write_back->app(
true, cont,
w.tuple({stored->proj(0), wb_buf}));
203 acc =
w.tuple({wb_mem,
init});
206 red_iters.reserve(rr);
207 for (
u64 j = 0; j < rr; ++j) {
208 auto dim = Sr->proj(nloops, ro + j);
210 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(
"forIn_" + std::to_string(j)));
211 auto [iter, new_acc, yield] = body->vars<3>();
215 current_mut->set(
true, for_call);
218 auto [red_mem, elem_acc] = acc->projs<2>();
220 DefVec iters_v = out_iters;
221 iters_v.insert(iters_v.end(), red_iters.begin(), red_iters.end());
222 auto iters =
w.tuple(iters_v);
226 DefVec input_elems(nis_nat);
227 for (
u64 i = 0; i < nis_nat; ++i) {
228 auto in_buf = new_inputs->proj(nis_nat, i);
229 auto [mc_mem, coords]
230 = affine_map(accs->proj(nis_nat, i), Ris->proj(nis_nat, i), n, Sr, Sis->proj(nis_nat, i), iters, cur);
233 auto [rd_mem, rd_val]
234 =
buffer::op_read(ir, is_, iT, cur, in_buf, fold_index(Sis->proj(nis_nat, i), coords))->
projs<2>();
236 input_elems[i] = rd_val;
240 current_mut->app(
true, comb,
w.tuple({w.tuple({cur, elem_acc, w.tuple(input_elems)}), cont}));
245const Def* LowerAff::lower_broadcast(
const App* app) {
246 auto&
w = new_world();
247 auto callee = app->callee()->as<
App>();
248 auto [s_in, s_out] = rewrite(callee->arg())->projs<2>();
249 auto [op_mem, input] = rewrite(app->arg())->projs<2>();
250 auto result_ty = rewrite(app->type());
252 auto r_nat = s_out->num_projs();
254 auto mem_ty =
w.call<
mem::M>(0);
255 auto fun =
w.mut_fun(
w.sigma({mem_ty, input->type()}), result_ty)->set(
"broadcast");
257 auto [fun_mem, in_buf] = fun->var(0_n)->projs<2>();
258 auto cont = fun->var(1);
264 const Def* acc =
w.tuple({a_mem, out_buf});
265 auto current_mut = fun;
267 out_iters.reserve(r_nat);
268 for (
size_t i = 0; i < r_nat; ++i) {
269 auto dim = s_out->proj(r_nat, i);
271 auto [body, for_call] = counting_for(bound, acc, cont,
w.sym(
"bcast_" + std::to_string(i)));
272 auto [iter, new_acc, yield] = body->vars<3>();
276 current_mut->set(
true, for_call);
279 auto [loop_mem, loop_buf] = acc->projs<2>();
282 auto iters =
w.tuple(out_iters);
283 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>();
284 auto [wr_mem, wr_buf]
285 =
buffer::op_write(out_r, out_s, out_T, rd_mem, loop_buf, fold_index(s_out, iters), rd_val)->
projs<2>();
286 current_mut->app(
true, cont,
w.tuple({wr_mem, loop_buf}));
291const Def* LowerAff::lower_pad(
const App* app) {
292 auto&
w = new_world();
293 auto c = rewrite(app->callee())->as<
App>();
297 auto [Tr, params] =
c->uncurry_args<2>();
298 auto [s_in, s_out,
mode, lo, hi] = params->projs<5>();
299 auto [op_mem, input, value] = rewrite(app->arg())->projs<3>();
300 auto result_ty = rewrite(app->type());
304 if (!r_l || !mode_l) {
305 WLOG(
"{} doesn't have a lowering-time known rank/mode", app);
306 return RWPhase::rewrite_imm_App(app);
309 auto mode_nat = *mode_l;
310 auto i64 =
w.type_i64();
313 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
315 auto compute = [&](
const DefVec& iters,
const Def* ins,
const Def* mem) -> std::pair<const Def*, const Def*> {
316 auto [in_buf, fill] = ins->projs<2>();
320 for (
u64 d = 0;
d < rn; ++
d) {
327 valid.push_back(v_d);
328 idx_i64 = sel(v_d, in_d,
w.lit_i64(0));
332 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({in_d, sin_m1}))}));
338 if (mode_nat != 0)
return {rd_mem, elem};
339 auto all_valid = valid.empty() ?
w.lit_tt() : valid[0];
340 for (
u64 d = 1;
d < valid.size(); ++
d)
342 return {rd_mem, sel(all_valid, elem, fill)};
345 return build_pointwise(w, result_ty, op_mem,
w.tuple({input, value}), s_out, rn,
"pad", compute);
348const Def* LowerAff::lower_concat(
const App* app) {
349 auto&
w = new_world();
350 auto c = rewrite(app->callee())->as<
App>();
354 auto [TnisR, ax, Sis, s_out] =
c->uncurry_args<4>();
355 auto [T, nis,
r] = TnisR->projs<3>();
356 auto [op_mem, op_is] = rewrite(app->arg())->projs<2>();
357 auto result_ty = rewrite(app->type());
362 if (!nis_l || !r_l || !ax_l) {
363 WLOG(
"{} doesn't have lowering-time known nis/rank/axis", app);
364 return RWPhase::rewrite_imm_App(app);
366 auto nisn = *nis_l, rn = *r_l, axn = *ax_l;
372 for (
u64 i = 0; i < nisn; ++i) {
373 off[i] =
w.lit_i64(acc_off);
376 WLOG(
"{} input {} has a non-literal extent along the concat axis", app, i);
377 return RWPhase::rewrite_imm_App(app);
383 auto sel = [&](
const Def*
cond,
const Def*
t,
const Def*
f) {
return w.extract(
w.tuple({f, t}), cond); };
385 auto compute = [&](
const DefVec& iters,
const Def* ins,
const Def* mem) -> std::pair<const Def*, const Def*> {
386 auto o_ax = iters[axn];
387 const Def* cur = mem;
389 auto read_i = [&](
u64 i) ->
const Def* {
390 auto in_buf = ins->proj(nisn, i);
392 auto Sis_i = Sis->proj(nisn, i);
393 auto e_i_m1 =
w.lit_i64(ext[i] - 1);
396 w.tuple({w.lit_i64(0), w.call(core::extrema::smin, w.tuple({loc, e_i_m1}))}));
398 for (
u64 d = 0;
d < rn; ++
d) {
399 auto idx_i64 = (
d == axn) ? clamp : iters[
d];
402 auto [rd_mem, rd_val]
408 auto result = read_i(0);
409 for (
u64 i = 1; i < nisn; ++i) {
411 result = sel(cond, read_i(i), result);
413 return {cur, result};
416 return build_pointwise(w, result_ty, op_mem, op_is, s_out, rn,
"concat", compute);
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 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.
bool is_bootstrapping() const
Returns whether we are currently bootstrapping (rewriting annexes).
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_cps2ds_dep(const Def *k)
Lam * mut_con(World &w, nat_t a=0)
Yields con[mem.M 0].
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 >