3#include <fe/worklist.h>
24bool contains_pi(
const Def* t) {
25 if (
t->isa<Pi>())
return true;
26 if (
auto sig =
t->isa_imm<Sigma>())
27 for (
auto op : sig->ops())
28 if (contains_pi(op))
return true;
36const Def* splat_scalar(
const Def* d) {
37 if (!
d->isa<Pack>())
return nullptr;
38 while (
auto pack =
d->isa<Pack>())
40 return (
d->is_closed() && !
d->type()->isa<Arr>()) ?
d :
nullptr;
43std::pair<Lam*, const Def*> counting_for(
const Def* bound,
const Def* acc,
const Def* exit, Sym name) {
44 auto&
w = bound->world();
45 auto body =
w.mut_con({
w.type_i64(), acc->type(),
w.cn(acc->type())})->
set(name);
46 auto for_loop =
w.call<
affine::For>(body, exit,
Defs{
w.lit_i64(0), bound,
w.lit_i64(1), acc});
47 return {body, for_loop};
51bool wants_fresh_mem(
const App* app) {
58bool is_tensor_op(
const App* app) {
64void LowerToMem::collect_tensor_types() {
67 auto gate = [](
const char* why,
const Def* culprit) { fe::throwf(
"cannot bufferize: {} (`{}`)", why, culprit); };
68 auto elem = [&gate](
const Def* T) {
69 if (T->isa<Arr>()) gate(
"tensor with array element type", T);
73 auto add_tensor_ty = [
this](
const Def*
t) {
74 if (
t->isa<Arr>()) tensor_ty_.emplace(t);
77 fe::BFSWorklist<DefSet> wl;
78 auto push = [&wl](
const Def*
d) {
82 for (
auto mut :
old_world().externals().muts())
88 if (
auto app = def->isa<App>()) {
89 if (
auto [axm, curry, trip] =
Axm::get(app); axm && curry == 0 && axm->plugin() == tensor::Plugin_Id)
96 add_tensor_ty(app->arg()->proj(n, 1)->type());
100 add_tensor_ty(app->type());
102 auto [meta, s_out] = app->callee()->as<
App>()->uncurry_args<2>();
103 auto [T,
r] = meta->projs<2>();
105 if (!
Lit::isa<u64>(r)) gate(
"non-literal rank of `tensor.generate`", app);
106 add_tensor_ty(app->type());
109 add_tensor_ty(app->type());
110 auto [nis_nps, meta, shapes, in_tys, comb_init, acc_out, accs]
111 = app->callee()->as<
App>()->uncurry_args<7>();
112 auto [nis, nps] = nis_nps->projs<2>();
113 auto [To, Tp, Ro, Rn, TSched] = meta->projs<5>();
114 auto [Tis, Ris, Sis, Tps, Rps, Sps] = in_tys->projs<6>();
115 auto [is, post_is] = app->arg()->projs<2>();
118 for (
u64 i = 0; i < *nis_l; ++i) {
119 elem(Tis->proj(*nis_l, i));
120 add_tensor_ty(is->proj(*nis_l, i)->type());
123 for (
u64 j = 0; j < *nps_l; ++j) {
124 elem(Tps->proj(*nps_l, j));
125 add_tensor_ty(post_is->proj(*nps_l, j)->type());
131 add_tensor_ty(app->type());
132 add_tensor_ty(app->arg()->proj(3, 2)->type());
135 auto [Tr, s_in, params] = app->callee()->as<
App>()->uncurry_args<3>();
136 auto [T,
r] = Tr->projs<2>();
137 auto [
mode, lo, hi] = params->projs<3>();
140 add_tensor_ty(app->type());
141 add_tensor_ty(app->arg()->proj(2, 0)->type());
144 auto [TnisR, ax, Sis] = app->callee()->as<
App>()->uncurry_args<3>();
145 auto [T, nis,
r] = TnisR->projs<3>();
150 if (!nis_l || !r_l || !ax_l) {
151 gate(
"non-literal nis/rank/axis of `tensor.concat`", app);
153 add_tensor_ty(app->type());
154 for (
u64 i = 0; i < *nis_l; ++i) {
157 gate(
"non-literal extent along the concat axis", app);
158 add_tensor_ty(app->arg()->proj(*nis_l, i)->type());
165 auto [Tr, shapes, dim] = app->callee()->as<
App>()->uncurry_args<3>();
166 auto [T,
r] = Tr->projs<2>();
169 add_tensor_ty(app->type());
170 add_tensor_ty(app->arg()->proj(2, 0)->type());
171 add_tensor_ty(app->arg()->proj(2, 1)->type());
173 auto [Tr, shapes, dim] = app->callee()->as<
App>()->uncurry_args<3>();
174 auto [T,
r] = Tr->projs<2>();
177 add_tensor_ty(app->type());
178 add_tensor_ty(app->arg()->proj(3, 0)->type());
179 add_tensor_ty(app->arg()->proj(3, 1)->type());
180 add_tensor_ty(app->arg()->proj(3, 2)->type());
181 }
else if (
auto [axm, curry, trip] =
Axm::get(app);
182 axm && curry == 0 && axm->plugin() == tensor::Plugin_Id) {
184 gate(
"unbufferizable tensor op", app);
190 if (is_tensor_op(app)) {
191 auto args = fe::BFSWorklist<DefSet>();
192 for (
const App* a = app;
a;
a =
a->callee()->isa<App>())
194 while (!
args.empty()) {
196 if (
auto k =
d->isa_mut<Lam>()) op_args_.emplace(k);
197 for (
auto op :
d->ops())
198 if (op)
args.push(op);
203 for (
auto op : def->ops())
208 for (
auto mut :
old_world().externals().muts())
209 if (
auto lam = mut->isa_mut<Lam>(); lam && is_tensor_fn(lam)) tensor_fns_.emplace(lam);
212 if (tensor_fns_.empty() && !ops_seen_)
return;
216 for (
auto old_fn : tensor_fns_) {
217 auto dom = old_fn->type()->dom();
218 auto n = dom->num_projs();
219 for (
size_t i = 0; i != n; ++i)
220 if (
auto pi =
Pi::isa_cn(dom->proj(n, i)); pi && contains_pi(pi->dom()))
221 return gate(
"higher-order bufferized function", old_fn);
225 fe::BFSWorklist<DefSet> wl2;
226 for (
auto mut :
old_world().externals().muts())
228 while (!wl2.empty()) {
229 auto def = wl2.pop();
231 for (
auto op : def->ops()) {
233 if (
auto fn =
op->isa_mut<Lam>()) {
236 if (tensor_fns_.contains(fn))
237 if (!(def->isa<App>() && def->as<App>()->callee() == op) && !def->isa<Var>())
238 return gate(
"bufferized function used as a value", fn);
239 if (!tensor_fns_.contains(fn) && !op_args_.contains(fn) && mentions_tensor(fn->type()->dom())) {
243 if (fn->is_external() || !
Pi::isa_cn(fn->type()))
244 return gate(
"unconvertible tensor-typed function", fn);
249 if (def->type()) wl2.push(def->type());
254 collect_tensor_types();
256 if (tensor_fns_.empty() && !ops_seen_)
return;
259 assert(pending_.empty());
262const Def* LowerToMem::buf_of(
const Def* arr_ty) {
266 while (
auto arr = cur->isa<
Arr>()) {
267 dims.push_back(
rewrite(arr->arity()));
273const Def* LowerToMem::fold_index(
const Def*
shape,
const Def* idx) {
275 auto r =
shape->num_projs();
277 for (
size_t i = 0; i != r; ++i)
278 if (
auto l =
Lit::isa<u64>(
shape->proj(r, i)); !(l && *l == 1)) out.push_back(idx->proj(r, i));
282const Def* LowerToMem::bot_mem() {
284 return w.bot(w.call<
mem::M>(0));
287const Def* LowerToMem::fresh_mem() {
289 pending_.push_back(k);
293void LowerToMem::wrap_fresh_mem(Lam* new_lam) {
295 auto filter = new_lam->filter();
296 auto body = new_lam->body();
299 for (
auto k : pending_ | std::views::reverse) {
301 body =
w.app(
w.annex<
mem::fresh>(),
w.tuple({w.lit_nat_0(), k}));
303 new_lam->set(filter, body);
311 if (
auto app = old_def->isa<
App>(); app && wants_fresh_mem(app)) {
312 if (
auto i = fresh_memo_.find(app); i != fresh_memo_.end())
return i->second;
314 fresh_memo_.emplace(app, new_def);
321bool LowerToMem::mentions_tensor(
const Def* t)
const {
322 if (tensor_ty_.contains(t))
return true;
323 if (
auto sig = t->isa<
Sigma>()) {
324 for (
auto op : sig->ops())
325 if (mentions_tensor(op))
return true;
326 }
else if (
auto pi = t->isa<
Pi>()) {
327 return mentions_tensor(pi->dom());
332bool LowerToMem::is_tensor_fn(
Lam* lam)
const {
336const Def* LowerToMem::conv_boundary(
const Def* t) {
337 if (tensor_ty_.contains(t))
return buf_of(t);
338 if (
auto sig = t->isa_imm<Sigma>(); sig && mentions_tensor(sig)) {
339 auto n = sig->num_ops();
341 for (
size_t i = 0; i != n; ++i)
342 ops[i] = conv_boundary(sig->op(i));
354 auto _ = fe::Restore(pending_, {});
355 auto __ = fe::Restore(fresh_memo_, {});
356 auto new_def = conv_mut_Lam(lam);
357 if (!pending_.empty()) wrap_fresh_mem(new_def->as_mut<
Lam>());
361const Def* LowerToMem::conv_mut_Lam(
Lam* lam) {
363 auto rebuild = [&](
const Def* dom) {
373 if (is_tensor_fn(lam)) {
377 DefVec doms(n, [&](
size_t i) ->
const Def* {
378 auto d = dom->proj(n, i);
379 if (
auto pi =
Pi::isa_cn(d))
return w.cn(conv_boundary(pi->dom()));
380 return conv_boundary(d);
382 return rebuild(w.sigma(doms));
389 if (
auto pi =
Pi::isa_cn(lam->
type()); pi && mentions_tensor(pi->dom()))
390 return rebuild(conv_boundary(pi->dom()));
392 return RWPhase::rewrite_mut_Lam(lam);
412 if (
auto callee = app->
callee()->
isa_mut<
Lam>(); callee && tensor_fns_.contains(callee))
413 return lower_call(app, callee);
420 auto elementwise = [&](
const Def* callee) {
421 if (
auto lam = callee->
isa_mut<
Lam>())
return op_args_.contains(lam);
422 for (
auto d = callee; d;) {
423 if (
auto ex = d->isa<
Extract>()) {
427 if (
auto var = d->isa<
Var>())
428 if (
auto lam = var->binder()->
isa_mut<
Lam>())
return op_args_.contains(lam);
433 if (elementwise(app->
callee()))
return RWPhase::rewrite_imm_App(app);
438 return RWPhase::rewrite_imm_App(app);
441const Def* LowerToMem::lower_call(
const App* app,
Lam* old_callee) {
443 auto new_callee =
rewrite(old_callee);
444 auto dom = old_callee->
type()->
dom();
448 for (
size_t i = 0; i != n; ++i) {
449 auto d = dom->proj(n, i);
450 auto a = app->
arg()->
proj(n, i);
455 return w.app(new_callee, w.tuple(
args));
458const Def* LowerToMem::splat_buffer(
const Def* arr_ty,
const Def* scalar) {
466const Def* LowerToMem::materialize(
const Def* old_ty,
const Def* old_arg) {
468 if (tensor_ty_.contains(old_ty)) {
471 if (
auto c = splat_scalar(old_arg))
return splat_buffer(old_ty,
rewrite(c));
478 if (
auto sig = old_ty->
isa_imm<Sigma>(); sig && mentions_tensor(sig)) {
479 auto n = sig->num_ops();
481 for (
size_t i = 0; i != n; ++i)
488const Def* LowerToMem::to_buffer(
const Def* val,
const Def* old) {
490 return materialize(old->type(), old);
493const Def* LowerToMem::buffer_list(
const Def* list,
const Def* old_list,
const Def* n) {
495 if (!n_l)
return list;
497 DefVec ins(*n_l, [&](
size_t i) {
return to_buffer(list->proj(*n_l, i), old_list->proj(*n_l, i)); });
503const Def* LowerToMem::lower_get(
const App* app) {
505 auto arg =
rewrite(app->arg());
506 auto [index, arr] = arg->projs<2>();
507 auto [T,
r,
s] =
c->args<3>();
508 arr = to_buffer(arr, app->arg()->proj(2, 1));
510 if (!buf)
return RWPhase::rewrite_imm_App(app);
511 auto [br, bs, bT] = buf->args<3>();
517const Def* LowerToMem::lower_set(
const App* app) {
519 auto arg =
rewrite(app->arg());
520 auto [index, arr, x] = arg->projs<3>();
521 auto [T,
r,
s] =
c->args<3>();
522 arr = to_buffer(arr, app->arg()->proj(3, 1));
524 if (!buf)
return RWPhase::rewrite_imm_App(app);
525 auto [br, bs, bT] = buf->args<3>();
526 auto fidx = fold_index(s, index);
528 if (reuse_in_place(app)) {
541const Def* LowerToMem::lower_splat(
const App* app) {
542 auto value = app->arg()->proj(2, 1);
547 if (!app->type()->isa<Arr>())
return rewrite(value);
548 return splat_buffer(app->type(),
rewrite(value));
551const Def* LowerToMem::lower_generate(
const App* app) {
553 auto [meta, s_out] =
rewrite(app->callee())->as<
App>()->uncurry_args<2>();
554 auto [T,
r] = meta->projs<2>();
556 if (!r_l)
return RWPhase::rewrite_imm_App(app);
558 auto body =
rewrite(app->arg());
562 if (!app->type()->isa<Arr>()) {
563 DefVec zeros(rn, [&](
size_t) {
return w.lit_i64(0); });
564 return w.call(body,
w.tuple(zeros));
567 auto out_ty = buf_of(app->type());
569 auto mem_ty =
w.call<
mem::M>(0);
570 auto result_ty =
w.sigma({mem_ty, out_ty});
571 auto unit =
w.tuple(
Defs{});
572 auto fun =
w.mut_fun(
w.sigma({mem_ty, unit->type()}), result_ty)->set(
"tensor_generate");
574 auto [
args, cont] = fun->vars<2>();
575 auto [fun_mem, ignored] =
args->projs<2>();
578 const Def* acc =
w.tuple({alloc_mem, out});
582 for (
u64 d = 0;
d < rn; ++
d) {
584 auto [loop, loop_call] = counting_for(bound, acc, cont,
w.sym(
"generate_" + std::to_string(d)));
585 auto [iter, next, yield] = loop->vars<3>();
586 iters.push_back(iter);
589 current->set(
true, loop_call);
593 auto [loop_mem, loop_out] = acc->projs<2>();
594 auto element =
w.call(body,
w.tuple(iters));
596 for (
u64 d = 0;
d < rn; ++
d)
598 auto [write_mem, written]
599 =
buffer::op_write(br, bs, bT, loop_mem, loop_out, fold_index(s_out,
w.tuple(coords)), element)->
projs<2>();
600 current->app(
true, cont,
w.tuple({write_mem, written}));
601 auto [call_mem, call_out] = call->projs<2>();
605const Def* LowerToMem::lower_broadcast(
const App* app) {
610 auto arg =
rewrite(app->arg());
611 auto [s_in, s_out, input] = arg->projs<3>();
612 auto [T,
r] =
c->args<2>();
615 if (s_in == s_out)
return input;
617 input = to_buffer(input, app->arg()->proj(3, 2));
622 if (!in_buf)
return splat_buffer(app->type(), input);
625 auto [bri, bsi, biT] = in_buf->args<3>();
631 op =
w.app(op,
w.tuple({T, bri, bsi, bro, bso, r}));
632 op =
w.app(op,
w.tuple({s_in, s_out}));
633 auto [m, out] =
w.app(op,
w.tuple({fresh_mem(), input}))->projs<2>();
637const Def* LowerToMem::lower_map_reduce(
const App* app) {
643 auto [is, post_is] =
args->projs<2>();
645 auto [nis_nps, meta, shapes, in_tys, comb_init, acc_out, accs] =
c->uncurry_args<7>();
646 auto [nis, nps] = nis_nps->projs<2>();
647 auto [comb,
init, post] = comb_init->projs<3>();
652 auto [old_is, old_post_is] = app->arg()->projs<2>();
653 is = buffer_list(is, old_is, nis);
654 post_is = buffer_list(post_is, old_post_is, nps);
655 if (!is || !post_is)
return RWPhase::rewrite_imm_App(app);
659 auto mem_ty =
w.call<
mem::M>(0);
660 auto inner = comb->type()->as<
Pi>()->dom()->proj(2, 0);
661 auto [cTo, ins_ty] = inner->projs<2>();
662 auto memcomb =
w.mut_fun(
w.sigma({mem_ty, cTo, ins_ty}),
w.sigma({mem_ty, cTo}))->set(
"memComb");
663 auto [cdata, cret] = memcomb->vars<2>();
664 auto [cm, cacc, cins] = cdata->projs<3>();
665 auto after =
w.mut_con(cTo)->set(
"afterComb");
666 after->app(
true, cret,
w.tuple({cm, after->var()}));
667 memcomb->set(
true,
w.app(comb,
w.tuple({w.tuple({cacc, cins}), after})));
671 auto pTp = meta->proj(5, 1);
672 auto pin = post->type()->as<
Pi>()->dom()->proj(2, 0);
673 auto [pTo, pexts_ty] = pin->projs<2>();
674 auto mempost =
w.mut_fun(
w.sigma({mem_ty, pTo, pexts_ty}),
w.sigma({mem_ty, pTp}))->set(
"memPost");
675 auto [pdata, pret] = mempost->vars<2>();
676 auto [pm, pacc, pexts] = pdata->projs<3>();
677 auto pafter =
w.mut_con(pTp)->set(
"afterPost");
678 pafter->app(
true, pret,
w.tuple({pm, pafter->var()}));
679 mempost->set(
true,
w.app(post,
w.tuple({
w.tuple({pacc, pexts}), pafter})));
685 op =
w.app(op, nis_nps);
686 op =
w.app(op, meta);
687 op =
w.app(op, shapes);
688 op =
w.app(op, in_tys);
689 op =
w.app(op,
w.tuple({memcomb, init, mempost}));
690 op =
w.app(op, acc_out);
691 op =
w.app(op, accs);
692 auto [m, out] =
w.app(op,
w.tuple({fresh_mem(), is, post_is}))->projs<2>();
696const Def* LowerToMem::lower_pad(
const App* app) {
699 auto&
w = new_world();
700 auto c = rewrite(app->callee())->as<
App>();
701 auto [input, value] = rewrite(app->arg())->projs<2>();
703 auto [Tr, s_in, params] =
c->uncurry_args<3>();
704 auto [T,
r] = Tr->projs<2>();
705 auto [
mode, lo, hi] = params->projs<3>();
707 input = to_buffer(input, app->arg()->proj(2, 0));
709 return RWPhase::rewrite_imm_App(app);
712 if (!r_l)
return RWPhase::rewrite_imm_App(app);
717 DefVec so(*r_l, [&](
size_t d) {
return add(
add(lo->proj(*r_l, d), s_in->proj(*r_l, d)), hi->proj(*r_l, d)); });
718 auto s_out =
w.tuple(so);
725const Def* LowerToMem::lower_concat(
const App* app) {
728 auto&
w = new_world();
729 auto c = rewrite(app->callee())->as<
App>();
730 auto arg = rewrite(app->arg());
732 auto [TnisR, ax, Sis] =
c->uncurry_args<3>();
733 auto [T, nis,
r] = TnisR->projs<3>();
738 if (!nis_l || !r_l || !ax_l)
return RWPhase::rewrite_imm_App(app);
740 auto inputs = buffer_list(arg, app->arg(), nis);
741 if (!inputs)
return RWPhase::rewrite_imm_App(app);
747 for (
u64 i = 0; i < *nis_l; ++i) {
749 if (!e)
return RWPhase::rewrite_imm_App(app);
752 DefVec so(*r_l, [&](
size_t d) {
return d == *ax_l ?
w.lit_nat(sum_ax) : Sis->proj(*nis_l, 0)->proj(*r_l, d); });
753 auto s_out =
w.tuple(so);
760 op =
w.app(op,
w.tuple({T, nis, r}));
763 op =
w.app(op, s_out);
764 auto [m, out] =
w.app(op,
w.tuple({fresh_mem(), inputs}))->projs<2>();
768const Def* LowerToMem::lower_gather(
const App* app) {
769 auto&
w = new_world();
770 auto c = rewrite(app->callee())->as<
App>();
772 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
773 auto [T,
r] = Tr->projs<2>();
774 auto [s_src, s_idx] = shapes->projs<2>();
776 auto [old_input, old_index] = app->args<2>();
777 auto make_buffer = [&](
const Def* old_arg,
const Def*
shape) {
778 auto value =
materialize(old_arg->type(), old_arg);
783 auto input = make_buffer(old_input, s_src);
784 auto index = make_buffer(old_index, s_idx);
787 auto [out_mem, out_buf] =
w.call(op,
Defs{s_src, s_idx}, dim,
Defs{fresh_mem(), input, index})->projs<2>();
788 if (app->type()->isa<Arr>())
return out_buf;
793const Def* LowerToMem::lower_scatter(
const App* app) {
794 auto&
w = new_world();
795 auto c = rewrite(app->callee())->as<
App>();
797 auto [Tr, shapes, dim] =
c->uncurry_args<3>();
798 auto [T,
r] = Tr->projs<2>();
799 auto [s_src, s_idx, s_updates] = shapes->projs<3>();
801 auto [old_input, old_index, old_updates] = app->args<3>();
802 auto make_buffer = [&](
const Def* old_arg,
const Def*
shape) {
803 auto value =
materialize(old_arg->type(), old_arg);
808 auto input = make_buffer(old_input, s_src);
809 auto index = make_buffer(old_index, s_idx);
810 auto updates = make_buffer(old_updates, s_updates);
813 auto [out_mem, out_buf]
814 =
w.call(op,
Defs{s_src, s_idx, s_updates}, dim,
Defs{fresh_mem(), input, index, updates})->projs<2>();
815 if (app->type()->isa<Arr>())
return out_buf;
const Def * callee() const
A (possibly paramterized) Array.
static auto isa(const Def *def)
static std::tuple< const Axm *, u8, u8 > get(const Def *def)
Yields currying counter of def.
const Def * proj(nat_t a, nat_t i) const
Similar to World::extract while assuming an arity of a, but also works on Sigmas and Arrays.
T * isa_mut() const
If this is mutable, it will cast constness away and perform a dynamic_cast to T.
DbgKey dbg_key() const
Cheap handle for other->set(this->dbg_key()).
const Def * var(nat_t a, nat_t i) noexcept
auto projs(F f) const
Splits this Def via Def::projections into an Array (if A == std::dynamic_extent) or std::array (other...
nat_t num_vars() noexcept
const Def * type() const noexcept
Yields the "raw" type of this Def (maybe nullptr).
bool is_external() const noexcept
nat_t num_projs() const
Yields Def::arity(), if it is a Lit, or 1 otherwise.
const T * isa_imm() const
const Def * filter() const
Lam * set(Filter filter, const Def *body)
static std::optional< T > isa(const Def *def)
const fe::Vector< std::string > & args()
Command-line arguments passed to this Phase's plugin via -X <plugin>:<arg>.
A dependent function type.
static const Pi * isa_cn(const Def *d)
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 *)
A variable introduced by a binder (mutable).
const Def * sigma(Defs ops)
const Def * tuple(Defs ops)
Lam * mut_con(const Def *dom)
const Def * rewrite_mut_Lam(Lam *) override
const Def * rewrite(const Def *) override
void start() override
Actual entry.
const Def * rewrite_imm_App(const App *) override
const Def * type_buf(const Def *r, const Def *s, const Def *T)
The buffer type buffer.Buf (r, s, T).
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_lit(const Def *r, const Def *s, const Def *T, const Def *mem, const Def *val)
buffer.lit (r, s, T) (mem, val) ↦ [mem.M 0, buffer.Buf (r, s, T)] (every element initialised to val).
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_init(const Def *r, const Def *s, const Def *T, const Def *mem, const Def *val)
buffer.init (r, s, T) (mem, val) ↦ [mem.M 0, buffer.Buf (r, s, T)] (initialised with the array value ...
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].
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.
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec