25bool contains_gpu_init(
const Def* def,
DefSet& seen) {
26 if (
auto [_, ins] = seen.emplace(def); !ins)
return false;
28 for (
auto d : def->deps())
29 if (contains_gpu_init(d, seen))
return true;
34std::pair<Lam*, const Def*> counting_for(
const Def* bound,
const Def* acc,
const Def* exit, Sym name) {
35 auto&
w = bound->world();
36 auto acc_ty = acc->type();
37 auto body =
w.mut_con({
w.type_i64(), acc_ty,
w.cn(acc_ty)})->
set(name);
38 auto for_loop =
w.call<
affine::For>(body, exit,
Defs{
w.lit_i64(0), bound,
w.lit_i64(1), acc});
39 return {body, for_loop};
43std::pair<const Def*, const Def*>
44affine_map(
const Def* f,
const Def* m,
const Def* n,
const Def* sin,
const Def* sout,
const Def* idxs,
const Def* mem) {
45 auto&
w = mem->world();
50 a =
w.app(a, mem->type()->as<App>()->arg());
51 return w.app(a, mem)->projs<2>();
54const Def* fold_index(
const Def* shape,
const Def* idx) {
56 auto r =
shape->num_projs();
59 for (
size_t i = 0; i !=
r; ++i)
63 out.push_back(
idx->proj(r, i));
67 if (!dropped)
return idx;
72const Def* op_lea_tuple(
const Def* ptr,
const Def* tuple) {
73 auto n = tuple->num_projs();
75 for (
size_t i = 0; i != n; ++i)
81Lam* rebuild_lam_global_mem(Lam* lam,
const Def* Tout, Sym name) {
82 auto&
w = lam->world();
86 if (lam->num_vars() == 2) {
87 auto [_, Tin, extra_ty] = lam->var(0)->type()->projs<3>();
88 new_lam =
w.mut_con(
Defs{
w.sigma({global_ty, Tin, extra_ty}),
w.cn({global_ty, Tout})})->
set(name);
90 auto Tin = lam->var(1)->type();
91 auto extra_ty = lam->var(2)->type();
92 new_lam =
w.mut_con(
Defs{global_ty, Tin, extra_ty,
w.cn({global_ty, Tout})})->
set(name);
94 new_lam->set(
true, lam->reduce_body(new_lam->var()));
99void apply_cps(World& w, Lam* mut,
const Def* f,
DefVec parts,
const Def* k) {
100 auto dom =
f->type()->as<
Pi>()->dom();
101 if (dom->num_projs() == parts.size() + 1) {
102 parts.emplace_back(k);
103 mut->app(
true, f, parts);
105 mut->app(
true, f,
Defs{
w.tuple(parts), k});
109Vector<nat_t> row_major_strides(
const Vector<nat_t>& dims) {
110 Vector<nat_t> strides(dims.size());
112 for (
auto i = dims.size(); i-- != 0;) {
119std::pair<const Def*, DefVec> unflatten_index(World& w,
const Def* flat,
const Vector<nat_t>& dims,
const Def* mem) {
120 auto strides = row_major_strides(dims);
121 DefVec coords(dims.size());
122 for (
size_t d = 0;
d != dims.size(); ++
d) {
128 return {mem, coords};
135InputDesc extract_input_desc(
nat_t n,
const Def* Rs,
const Def* Ss,
const Def* Ts,
const Def* accs) {
137 for (
nat_t i = 0; i != n; ++i) {
138 desc.rs[i] = Rs->proj(n, i);
139 desc.ss[i] = Ss->proj(n, i);
140 desc.ts[i] = Ts->proj(n, i);
141 desc.accs[i] = accs->proj(n, i);
152Inputs alloc_copy_inputs(World& w,
const Def* m0,
const Def* m1,
Defs ris,
Defs sis,
Defs tis,
const Def* inputs) {
154 for (
size_t i = 0; i != ris.size(); ++i) {
156 {m0, m1, inputs->proj(ris.size(), i)});
162 return {m0, m1, dptrs};
165std::pair<const Def*, const Def*> alloc_output(World& w,
const Def* m1,
const Def* elem_ty,
const Def* So,
nat_t ro) {
166 auto arr_ty = elem_ty;
167 for (
auto d = ro;
d-- != 0;)
168 arr_ty =
w.arr(So->proj(ro, d), arr_ty);
173 nat_t n_groups, n_items, total;
176Grid grid_layout(
const Vector<nat_t>& out_dims) {
178 for (
auto d : out_dims)
180 nat_t n_items = std::min<nat_t>(total, 1024);
181 nat_t n_groups = (total + n_items - 1) / n_items;
182 return {n_groups, n_items, total};
187 DefVec rs, ss, dptrs, accs;
188 nat_t n()
const {
return dptrs.size(); }
192Lam* build_kernel(World& w,
195 const Vector<nat_t>& out_dims,
203 const Mapped& post_ins,
209 auto nps = post_ins.n();
210 auto ro = out_dims.size();
211 auto nloops_nat = ro + rr;
212 auto n =
w.lit_nat(nloops_nat);
219 DefVec arg_tys(nis + nps + 1);
220 for (
size_t i = 0; i != nis; ++i)
221 arg_tys[i] = ins.dptrs[i]->type();
222 for (
size_t j = 0; j != nps; ++j)
223 arg_tys[nis + j] = post_ins.dptrs[j]->type();
224 arg_tys[nis + nps] = out_dptr->type();
227 =
w.mut_con(
Defs{global_ty, shared_ty, const_ty, local_ty,
w.type_idx(grid.n_groups),
w.type_idx(grid.n_items),
228 w.sigma(
Defs{}),
w.sigma(arg_tys),
w.cn({global_ty, shared_ty, const_ty, local_ty})})
229 ->
set(
"mapReduceKernel");
230 auto [k_global, k_shared, k_const, k_local, group_id, item_id, k_shared_ptrs, k_args, k_ret] = kernel->vars<9>();
233 for (
size_t i = 0; i != nis; ++i)
234 k_dptrs[i] = k_args->proj(nis + nps + 1, i);
236 for (
size_t j = 0; j != nps; ++j)
237 k_post_dptrs[j] = k_args->proj(nis + nps + 1, nis + j);
238 auto k_out_dptr = k_args->proj(nis + nps + 1, nis + nps);
240 auto group_i64 = grid.n_groups == 1 ?
w.lit_i64(0) :
w.call(
core::conv::u,
w.lit_nat_0(), group_id);
241 auto item_i64 = grid.n_items == 1 ?
w.lit_i64(0) :
w.call(
core::conv::u,
w.lit_nat_0(), item_id);
247 auto early_return =
w.mut_con(
w.sigma(
Defs{}))->set(
"outOfRange");
248 early_return->app(
true, k_ret,
Defs{k_global, k_shared, k_const, k_local});
249 auto body =
w.mut_con(
w.sigma(
Defs{}))->set(
"inRange");
250 kernel->set(
true,
w.app(
w.extract(
w.tuple({early_return, body}), in_range),
w.tuple()));
252 auto write_back =
w.mut_con(
Defs{global_ty, To})->
set(
"writeBack");
253 auto [wb_mem, acc_final] = write_back->vars<2>();
254 auto [wb_mem2, wb_coords] = unflatten_index(w, flat, out_dims, wb_mem);
255 DefVec wb_idx = wb_coords;
256 for (
size_t j = 0; j != rr; ++j)
257 wb_idx.push_back(
w.call(
core::conv::u, Sr->proj(nloops_nat, ro + j),
w.lit_i64(0)));
258 auto [wc_mem, write_coords] = affine_map(acc_out, Ro, n, Sr, So,
w.tuple(wb_idx), wb_mem2);
262 for (
size_t j = 0; j != nps; ++j) {
263 auto [pc_mem, pcoords]
264 = affine_map(post_ins.accs[j], post_ins.rs[j], Ro, So, post_ins.ss[j], write_coords, pcur);
266 auto [rd_mem, rd_val]
267 =
w.call<
mem::load>(
Defs{pcur, op_lea_tuple(k_post_dptrs[j], fold_index(post_ins.ss[j], pcoords))})
270 post_elems[j] = rd_val;
273 auto after_post =
w.mut_con(
Defs{global_ty, Tp})->
set(
"afterPost");
274 auto [post_mem, elem_post] = after_post->vars<2>();
276 =
w.call<
mem::store>(
Defs{post_mem, op_lea_tuple(k_out_dptr, fold_index(So, write_coords)), elem_post});
277 after_post->app(
true, k_ret,
Defs{final_mem, k_shared, k_const, k_local});
278 apply_cps(w, write_back, global_post, {pcur, acc_final,
w.tuple(post_elems)}, after_post);
280 const Def* acc =
w.tuple({k_global,
init});
281 const Def* cont = write_back;
282 Lam* current_mut = body;
284 red_iters.reserve(rr);
285 for (
size_t j = 0; j != rr; ++j) {
286 auto dim = Sr->proj(nloops_nat, ro + j);
288 auto [rbody, for_call] = counting_for(bound, acc, cont,
w.sym(
"forRed_" + std::to_string(j)));
289 auto [iter, new_acc, yield] = rbody->vars<3>();
293 current_mut->set(
true, for_call);
296 auto [red_mem, elem_acc] = acc->projs<2>();
298 auto [body_mem, body_coords] = unflatten_index(w, flat, out_dims, red_mem);
299 DefVec iters_v = body_coords;
300 iters_v.insert(iters_v.end(), red_iters.begin(), red_iters.end());
301 auto iters =
w.tuple(iters_v);
305 for (
size_t i = 0; i != nis; ++i) {
306 auto [mc_mem, coords] = affine_map(ins.accs[i], ins.rs[i], n, Sr, ins.ss[i], iters, cur);
308 auto [rd_mem, rd_val]
309 =
w.call<
mem::load>(
Defs{cur, op_lea_tuple(k_dptrs[i], fold_index(ins.ss[i], coords))})->projs<2>();
311 input_elems[i] = rd_val;
314 apply_cps(w, current_mut, global_comb, {cur, elem_acc,
w.tuple(input_elems)}, cont);
319Lam* build_teardown(World& w,
329 auto mem_ty =
w.call<
mem::M>(0);
331 auto after_launch =
w.mut_con(
Defs{mem_ty, global_ty, const_ty})->
set(
"afterLaunch");
332 auto [post_mem, post_global, post_const] = after_launch->vars<3>();
337 auto [cb_mem, cb_global] = copy_back->projs<2>();
339 auto cur_global = cb_global;
340 for (
auto dptr : dptrs)
342 for (
auto dptr : post_dptrs)
347 after_launch->app(
true, cont,
Defs{final_mem, host_buf});
356 = std::ranges::any_of(
old_world().roots(), [&](
auto def) {
return contains_gpu_init(def, seen); });
358 log().w(
"not lowering any map-reduce operations to GPU: the program already contains an explicit `gpu.init`");
366 return Super::rewrite_imm_App(app);
369const Def* LowerMapReduce::lower_map_reduce_post(
const App* app) {
375 auto [nis_nps, meta, shapes, in_tys, comb_init, acc_out, accs_all] = c->uncurry_args<7>();
376 auto [nis, nps] = nis_nps->
projs<2>([](
auto d) {
return Lit::isa(d); });
377 auto [To, Tp, Ro, Rn, sched_ty] = meta->
projs<5>();
378 auto [So, Sr, sched] = shapes->
projs<3>();
379 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->
projs<6>();
380 auto [comb,
init, post] = comb_init->
projs<3>();
381 auto [accs, post_accs] = accs_all->projs<2>();
386 if (!nis || !nps || !ro_l || !rn_l || *rn_l < *ro_l) {
387 log().w(
"{} doesn't have lowering-time known rank counts (nis/nps/Ro/Rn)", app);
388 return Super::rewrite_imm_App(app);
393 auto rr = *rn_l - *ro_l;
395 Vector<nat_t> out_dims(ro);
397 for (
nat_t d = 0; d != ro; ++d) {
400 log().w(
"{} doesn't have a lowering-time known output (grid) shape", app);
401 return Super::rewrite_imm_App(app);
406 if (out_total == 0) {
407 log().w(
"{} has a zero-sized output, skipping GPU lowering", app);
408 return Super::rewrite_imm_App(app);
411 auto comb_lam = comb->isa_mut<
Lam>();
412 auto post_lam = post->isa_mut<
Lam>();
413 if (!comb_lam || !post_lam) {
414 log().w(
"{} doesn't have a lowering-time known combiner/epilogue", app);
415 return Super::rewrite_imm_App(app);
418 auto mem_ty =
w.call<
mem::M>(0);
420 auto [_, rewritten_inputs, rewritten_post_ins] = rewritten_arg->projs<3>();
421 auto fun =
w.mut_fun(
w.sigma({mem_ty, rewritten_inputs->type(), rewritten_post_ins->type()}), result_ty)
422 ->set(
"mapReduceAffGpu");
424 auto [fun_mem, new_inputs, new_post_ins] = fun->var(0_n)->projs<3>();
425 auto cont = fun->var(1);
427 auto [h_mem, h_global, h_const] =
w.app(
w.annex<
gpu::auto_init>(), fun_mem)->projs<3>();
429 auto in_desc = extract_input_desc(nis_n, Ris, Sis, Tis, accs);
430 auto inputs = alloc_copy_inputs(w, h_mem, h_global, in_desc.rs, in_desc.ss, in_desc.ts, new_inputs);
432 auto post_desc = extract_input_desc(nps_n, Rps, Sps, Tps, post_accs);
434 = alloc_copy_inputs(w, inputs.mem, inputs.global, post_desc.rs, post_desc.ss, post_desc.ts, new_post_ins);
436 auto [out_global, out_dptr] = alloc_output(w, post_inputs.global, Tp, So, ro);
438 auto global_comb = rebuild_lam_global_mem(comb_lam, To,
w.sym(
"combGlobal"));
439 auto global_post = rebuild_lam_global_mem(post_lam, Tp,
w.sym(
"postGlobal"));
441 auto grid = grid_layout(out_dims);
443 auto kernel = build_kernel(w, Ro, rr, out_dims, Sr, So, Mapped{in_desc.rs, in_desc.ss, inputs.dptrs, in_desc.accs},
444 To, acc_out,
init, global_comb,
445 Mapped{post_desc.rs, post_desc.ss, post_inputs.dptrs, post_desc.accs}, global_post, Tp,
448 DefVec kernel_arg_tys(nis_n + nps_n + 1);
449 for (
nat_t i = 0; i != nis_n; ++i)
450 kernel_arg_tys[i] = inputs.dptrs[i]->type();
451 for (
nat_t j = 0; j != nps_n; ++j)
452 kernel_arg_tys[nis_n + j] = post_inputs.dptrs[j]->type();
453 kernel_arg_tys[nis_n + nps_n] = out_dptr->type();
457 w.lit_ff(),
w.tuple()});
460 DefVec kernel_args = inputs.dptrs;
461 kernel_args.insert(kernel_args.end(), post_inputs.dptrs.begin(), post_inputs.dptrs.end());
462 kernel_args.push_back(out_dptr);
465 auto after_launch = build_teardown(w, Ro, So, Tp, inputs.dptrs, post_inputs.dptrs, out_dptr, cont);
466 auto launch_call =
w.app(
launch,
Defs{
w.tuple({post_inputs.mem, out_global, h_const}), after_launch});
467 fun->set(
true, launch_call);
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).
static std::optional< T > isa(const Def *def)
const fe::Log & log() const
bool is_bootstrapping() const
Returns whether we are currently bootstrapping (rewriting annexes).
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 *)
void start() final
Skips the whole phase if the program already contains an explicit gpu.init.
const Def * rewrite_imm_App(const App *) final
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)
const Def * op_lea(const Def *ptr, const Def *index)
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec
GIDSet< const Def * > DefSet