33 if (!map_reduce_ax)
return RWPhase::rewrite_imm_App(app);
38 auto [
mem, zero, comb, inputs] = map_reduce_ax->args<4>();
39 auto [n, S, T, m, NI, TI, SI] = map_reduce_ax->callee()->as<
App>()->
args<7>();
55 absl::flat_hash_map<u64, const Def*> dims;
56 absl::flat_hash_map<u64, const Def*> iterator;
62 auto n_lit = n->isa<
Lit>();
63 auto m_lit = m->
isa<
Lit>();
64 if (!n_lit || !m_lit)
return RWPhase::rewrite_imm_App(app);
66 auto n_nat = n_lit->get<
u64>();
67 auto m_nat = m_lit->
get<
u64>();
70 for (
u64 i = 0; i < n_nat; ++i)
71 dims[i] = S->proj(n_nat, i);
74 for (
u64 i = 0; i < m_nat; ++i) {
75 auto ni_lit =
Lit::isa(NI->proj(m_nat, i));
76 if (!ni_lit)
return RWPhase::rewrite_imm_App(app);
78 auto SI_i = SI->proj(m_nat, i);
79 input_dims.emplace_back(
DefVec(ni_nat, [&](
u64 j) {
return SI_i->proj(ni_nat, j); }));
80 n_input.push_back(ni_nat);
84 for (
u64 i = 0; i < m_nat; ++i) {
85 auto [indices, mat] = inputs->proj(m_nat, i)->projs<2>();
86 for (
u64 j = 0; j < n_input[i]; ++j) {
87 auto idx_lit =
Lit::isa(indices->proj(n_input[i], j));
88 if (!idx_lit)
return RWPhase::rewrite_imm_App(app);
89 u64 idx_nat = *idx_lit;
90 auto dim = input_dims[i][j];
91 if (!dims.contains(idx_nat)) {
93 }
else if (
auto dim_lit = dim->isa<
Lit>()) {
94 if (
auto prev_lit = dims[idx_nat]->isa<Lit>())
95 assert(dim_lit->get<
u64>() == prev_lit->get<
u64>() &&
"dimensions must be equal");
102 for (
auto [idx, dim] : dims)
103 (idx < n_nat ? out_indices : in_indices).push_back(idx);
105 std::sort(out_indices.begin(), out_indices.end());
106 std::sort(in_indices.begin(), in_indices.end());
117 for (
auto& [_, dim] : dims)
121 auto fun = w.mut_fun(w.call<
mem::M>(0),
rewrite(map_reduce_ax->type()))->set(
"mapRed");
141 auto cont = fun->var(1);
142 auto current_mut = fun;
145 DefVec acc = {current_mem, init_mat};
146 for (
auto idx : out_indices) {
147 auto dim_nat_def = dims[idx];
150 auto [body, for_call] = counting_for(dim, acc, cont, w.sym(
"forIn_"s + std::to_string(idx)));
151 auto [iter, new_acc, yield] = body->vars<3>();
152 auto [new_mem, new_acc_val] = new_acc->projs<2>();
154 iterator[idx] = w.call<
core::bitcast>(w.type_idx(dim_nat_def), iter);
155 acc = {new_mem, new_acc_val};
156 current_mut->set(
true, for_call);
161 auto elem_acc = zero->set(
"acc");
162 current_mem = acc[0];
163 auto wb_matrix = acc[1];
168 auto [wb_mem, elem_final] = write_back->
vars<2>();
170 auto output_it_tuple = w.tuple(
DefVec((
size_t)n_nat, [&](
u64 i) {
171 auto idx = out_indices[i];
172 if (idx != i)
ELOG(
"output indices must be consecutive 0..n-1 but {} != {}", idx, i);
173 assert(idx == i &&
"output indices must be consecutive 0..n-1");
174 return iterator[idx];
176 auto [wb_mem2, written_matrix]
178 write_back->app(
true, cont, {wb_mem2, written_matrix});
181 acc = {current_mem, elem_acc};
183 for (
auto idx : in_indices) {
184 auto dim_nat_def = dims[idx];
187 auto [body, for_call] = counting_for(dim, acc, cont, w.sym(
"forIn_"s + std::to_string(idx)));
188 auto [iter, new_acc, yield] = body->vars<3>();
189 auto [new_mem, new_acc_val] = new_acc->projs<2>();
191 iterator[idx] = w.call<
core::bitcast>(w.type_idx(dim_nat_def), iter);
192 acc = {new_mem, new_acc_val};
193 current_mut->set(
true, for_call);
196 current_mem = acc[0];
200 DefVec input_elems((
size_t)m_nat);
201 for (
u64 i = 0; i < m_nat; ++i) {
202 auto [input_idx_tup, input_matrix] = inputs->proj(m_nat, i)->projs<2>();
203 auto indices = input_idx_tup->projs(n_input[i]);
205 = w.tuple(
DefVec(n_input[i], [&](
u64 j) {
return iterator[indices[j]->as<
Lit>()->
get<u64>()]; }));
207 auto [new_mem, elem_i] =
op_read(current_mem,
rewrite(input_matrix), input_it_tuple)->
projs<2>();
208 current_mem = new_mem;
209 input_elems[i] = elem_i;
212 current_mut->app(
true, comb, {w.tuple({current_mem, elem_acc, w.tuple(input_elems)}), cont});
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)].