25constexpr u64 Loud_max_dispatch = 8;
28constexpr u64 Max_vec = 1024;
36bool is_kvec(
const Def* mat) {
38 if (!app)
return false;
40 auto perm = app->callee()->as<
App>()->callee()->as<
App>()->arg();
43 auto [p0, p1] =
perm->projs<2>();
49bool kvec_split(
Defs mats,
u64 s,
u64 j) {
return s + 1 == j && is_kvec(mats[j]); }
53using Mono = std::array<const Def*, 3>;
58 void add(Mono mono,
u64 coeff) {
59 for (
auto& [m, c] : terms_)
64 terms_.emplace_back(mono, coeff);
67 void add(
const Poly& other) {
68 for (
const auto& [m, c] : other.terms_)
72 u64 coeff(Mono mono)
const {
73 for (
const auto& [m, c] : terms_)
74 if (m == mono)
return c;
81 bool dominates(
const Poly& other)
const {
82 for (
const auto& [m, c] : terms_)
83 if (c > other.coeff(m))
return false;
87 std::string str()
const {
88 if (terms_.empty())
return "0";
89 auto s = std::string();
90 for (
const auto& [mono, coeff] : terms_) {
91 if (!
s.empty())
s +=
" + ";
92 s += std::format(
"{}", coeff);
94 if (d)
s += std::format(
"·{}", d);
100 fe::Vector<std::pair<Mono, u64>> terms_;
107Poly mul_cost(
const Def* x,
const Def* y,
const Def* z,
u64 vec,
bool kvec) {
110 auto extent = [&](
const Def*
d,
u64 pad) {
112 coeff *= lanes(*l,
pad);
114 syms.emplace_back(d);
118 extent(y, kvec ? vec : 1);
119 extent(z, kvec ? 1 : vec);
120 std::ranges::sort(syms, [](
const Def* a,
const Def* b) {
return a->gid() < b->gid(); });
123 std::ranges::copy(syms, mono.begin());
126 poly.add(mono, coeff);
130Poly cost_of(fe::View<Split> splits,
Defs mats,
Defs dims,
u64 vec) {
132 for (
auto [i, s, j] : splits)
133 poly.add(mul_cost(dims[i], dims[s + 1], dims[j + 1], vec, kvec_split(mats, s, j)));
138fe::Vector<Splits> bracketings(
u64 lo,
u64 hi) {
139 if (lo == hi)
return {
Splits()};
141 auto res = fe::Vector<Splits>();
142 for (
auto s = lo;
s != hi; ++
s)
143 for (
const auto& l : bracketings(lo, s))
144 for (
const auto& r : bracketings(s + 1, hi)) {
147 b.emplace_back(
Split{lo,
s, hi});
148 res.emplace_back(std::move(b));
156fe::Vector<Splits> pareto(fe::View<Splits> cands,
Defs mats,
Defs dims,
u64 vec) {
157 auto keep = fe::Vector<Splits>();
158 auto costs = fe::Vector<Poly>();
160 for (
const auto& cand : cands) {
161 auto cost = cost_of(cand, mats, dims, vec);
162 if (std::ranges::any_of(costs, [&](
const Poly& k) {
return k.dominates(cost); }))
continue;
163 for (
auto i = costs.size(); i-- != 0;)
164 if (cost.dominates(costs[i])) keep.erase(keep.begin() + i), costs.erase(costs.begin() + i);
165 keep.emplace_back(cand);
166 costs.emplace_back(std::move(cost));
172fe::Vector<u64> split_table(
const Splits& splits,
u64 n) {
173 auto table = fe::Vector<u64>(n * n, 0);
174 for (
auto [i, s, j] : splits)
175 table[i * n + j] =
s;
183std::optional<std::pair<fe::Vector<u64>, Poly>> matrix_chain_order(
Defs mats,
Defs dims,
u64 vec) {
184 auto n = dims.size() - 1;
185 auto cost = fe::Vector<Poly>(n * n);
186 auto split = fe::Vector<u64>(n * n, 0);
188 for (
auto len = 2_u64;
len <= n; ++
len) {
189 for (
auto i = 0_u64; i +
len <= n; ++i) {
190 auto j = i +
len - 1;
191 auto cands = fe::Vector<Poly>();
192 for (
auto s = i;
s != j; ++
s) {
193 auto c = mul_cost(dims[i], dims[s + 1], dims[j + 1], vec, kvec_split(mats, s, j));
194 c.add(cost[i * n + s]);
195 c.add(cost[(s + 1) * n + j]);
196 cands.emplace_back(std::move(c));
199 auto best = std::ranges::find_if(cands, [&](
const Poly& a) {
200 return std::ranges::all_of(cands, [&](
const Poly& b) {
return a.dominates(b); });
202 if (best == cands.end())
return {};
204 split[i * n + j] = i + (best - cands.begin());
205 cost[i * n + j] = std::move(*best);
209 return std::pair{std::move(split), std::move(cost[n - 1])};
215 auto num = [
this](
const char* key) -> std::optional<u64> {
219 auto end = val->data() + val->size();
220 if (
auto [ptr, ec] = std::from_chars(val->data(), end, n); ec != std::errc() || ptr != end) {
221 log().w(
"ignoring `-X tensor:{}={}`: not a number", key, *val);
227 if (
auto n = num(
"reassoc-max")) {
229 log().d(
"dispatch chains of up to {} matrices", *n);
230 if (*n > Loud_max_dispatch)
231 log().w(
"`-X tensor:reassoc-max={}` enumerates up to Catalan({}) bracketings", *n, *n - 1);
234 if (
auto n = num(
"reassoc-vec")) {
235 if (*n == 0 || *n > Max_vec) {
236 log().w(
"ignoring `-X tensor:reassoc-vec={}`: not between 1 and {} lanes", *n, Max_vec);
239 log().d(
"charge the vector loop in units of {} lanes", vec_);
250std::optional<Reassoc::Link> Reassoc::isa_link(
const Def* def,
const Def* ring)
const {
255 auto groups = app->callee()->as<
App>();
256 if (groups->callee()->as<
App>()->
arg() != ring)
return {};
258 auto [m, k, l] = groups->args<3>();
259 return Link{app, m, k, l};
262void Reassoc::flatten(
const Def* def,
const Def* ring,
const Def* rows,
DefVec& mats,
DefVec& dims,
Splits& orig) {
263 auto lo = mats.size();
265 if (
auto i = consumers_.find(def); i != consumers_.end() && i->second == 1)
266 if (
auto link = isa_link(def, ring)) {
267 auto [t1, t2] = link->app->args<2>();
268 flatten(t1, ring, link->m, mats, dims, orig);
269 auto mid = mats.size();
270 flatten(t2, ring, link->k, mats, dims, orig);
271 orig.emplace_back(Split{lo, mid - 1, mats.size() - 1});
275 mats.emplace_back(def);
276 dims.emplace_back(rows);
279const Def* Reassoc::build(
const Def* head,
Defs mats,
Defs dims, fe::View<u64> split,
u64 i,
u64 j) {
280 if (i == j)
return rewrite(mats[i]);
283 auto s = split[i * mats.size() + j];
284 auto t1 = build(head, mats, dims, split, i, s);
285 auto t2 = build(head, mats, dims, split, s + 1, j);
287 return w.app(
w.app(head, mkl), {t1, t2});
290const Def* Reassoc::cost_expr(
Defs mats,
Defs dims,
const Splits& splits) {
292 const Def*
sum =
nullptr;
295 auto extent = [&](
const Def*
d,
u64 pad) ->
const Def* {
300 for (
auto [i, s, j] : splits) {
301 auto kvec = kvec_split(mats, s, j);
302 auto p =
w.app(
w.annex(
core::nat::mul), {rewrite(dims[i]), extent(dims[s + 1], kvec ? vec_ : 1)});
303 p =
w.app(
w.annex(
core::nat::mul), {p, extent(dims[j + 1], kvec ? 1 : vec_)});
310const Def* Reassoc::dispatch(
const Def* head,
const Def* res_ty,
Defs mats,
Defs dims, fe::View<Splits> cands) {
312 auto n = mats.size();
313 auto pi =
w.pi(
w.sigma(), res_ty);
318 const Def* best =
nullptr;
319 const Def* best_cost =
nullptr;
320 for (
const auto& cand : cands) {
321 auto thunk =
w.mut_lam(pi)->set(
true, build(head, mats, dims, split_table(cand, n), 0, n - 1));
322 auto cost = cost_expr(mats, dims, cand);
325 best = thunk, best_cost = cost;
328 best =
w.extract(
w.tuple({best, (const Def*)thunk}), cheaper);
329 best_cost =
w.extract(
w.tuple({best_cost, cost}), cheaper);
333 return w.app(best,
w.tuple());
336const Def* Reassoc::reassoc(
const App* app) {
337 auto head = app->callee()->as<
App>()->callee()->as<
App>();
338 auto link = isa_link(app,
head->arg());
339 if (!link)
return nullptr;
341 auto [t1, t2] = app->args<2>();
345 flatten(t1,
head->arg(), link->m, mats, dims, orig);
346 auto mid = mats.size();
347 flatten(t2,
head->arg(), link->k, mats, dims, orig);
348 dims.emplace_back(link->l);
349 orig.emplace_back(Split{0_u64, mid - 1, mats.size() - 1});
352 auto n = mats.size();
353 if (n < 3)
return nullptr;
355 auto orig_cost = cost_of(orig, mats, dims, vec_);
357 if (n <= max_dispatch_) {
358 auto cands = pareto(bracketings(0, n - 1), mats, dims, vec_);
359 if (cands.size() != 1) {
360 log().d(
"dispatch chain {} over {} bracketings, written as {}", fe::Join(dims,
"×"), cands.size(),
362 return dispatch(
rewrite(head),
rewrite(app->type()), mats, dims, cands);
365 auto cost = cost_of(cands.front(), mats, dims, vec_);
366 if (orig_cost.dominates(cost))
return nullptr;
367 log().d(
"reassociate chain {}: {} → {} lane slots", fe::Join(dims,
"×"), orig_cost.str(), cost.str());
368 return build(
rewrite(head), mats, dims, split_table(cands.front(), n), 0, n - 1);
371 auto order = matrix_chain_order(mats, dims, vec_);
372 if (!order)
return nullptr;
374 auto& [split, cost] = *order;
375 if (orig_cost.dominates(cost))
return nullptr;
377 log().d(
"reassociate chain {}: {} → {} lane slots", fe::Join(dims,
"×"), orig_cost.str(), cost.str());
378 return build(
rewrite(head), mats, dims, split, 0, n - 1);
383 if (
auto res = reassoc(app))
return res;
384 return RWPhase::rewrite_imm_App(app);
static auto isa(const Def *def)
static std::optional< T > isa(const Def *def)
const fe::Log & log() const
const fe::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.
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 *)
const Def * rewrite_imm_App(const App *) final
void start() override
Actual entry.
fe::Vector< Split > Splits
A bracketing of a matrix chain, innermost node first.
One node of a bracketing: i … j splits after s.
DefMap< u64 > count_consumers(const World &world, Pred pred)
Counts the consumers of every def of world matched by pred.
fe::View< const Def * > Defs
fe::Vector< const Def * > DefVec
std::optional< std::string_view > arg_value(fe::View< std::string > args, Keys... keys)
Value of <key>=<value>; std::nullopt if none of keys carries one.