14#include <unordered_map>
15#include <unordered_set>
35constexpr double NaN = std::numeric_limits<double>::quiet_NaN();
58const std::unordered_set<gate_t> *g_shared_base_rvs =
nullptr;
60inline bool is_shared_base_rv(
gate_t rv)
62 return g_shared_base_rvs !=
nullptr && g_shared_base_rvs->count(rv) != 0;
70void collect_reachable_base_rvs(
const GenericCircuit &gc,
gate_t start,
71 std::unordered_set<gate_t> &out)
73 std::unordered_set<gate_t> seen;
74 std::stack<gate_t> stk;
76 while (!stk.empty()) {
77 gate_t g = stk.top(); stk.pop();
78 if (!seen.insert(g).second)
continue;
91 if (mw.size() == 3) { stk.push(mw[1]); stk.push(mw[2]); }
111double try_eval_constant(
const GenericCircuit &gc,
gate_t g)
116 catch (
const CircuitException &) {
return NaN; }
122 if (wires.empty())
return NaN;
124 double first = try_eval_constant(gc, wires[0]);
125 if (std::isnan(first))
return NaN;
129 if (wires.size() == 1)
return std::round(first);
131 double d = try_eval_constant(gc, wires[1]);
132 if (std::isnan(d))
return NaN;
133 const double f = std::pow(10.0, d);
134 return std::round(first * f) / f;
137 return wires.size() == 1 ? std::floor(first) : NaN;
139 return wires.size() == 1 ? std::ceil(first) : NaN;
141 return wires.size() == 1 ? std::fabs(first) : NaN;
145 return wires.size() == 1 ? first : NaN;
147 return wires.size() == 1
148 ?
static_cast<double>(
static_cast<float>(first)) : NaN;
151 for (std::size_t i = 1; i < wires.size(); ++i) {
152 double v = try_eval_constant(gc, wires[i]);
153 if (std::isnan(v))
return NaN;
160 for (std::size_t i = 1; i < wires.size(); ++i) {
161 double v = try_eval_constant(gc, wires[i]);
162 if (std::isnan(v))
return NaN;
168 if (wires.size() != 2)
return NaN;
169 double v = try_eval_constant(gc, wires[1]);
170 if (std::isnan(v))
return NaN;
174 if (wires.size() != 2)
return NaN;
175 double v = try_eval_constant(gc, wires[1]);
179 if (std::isnan(v) || v == 0.0)
return NaN;
183 if (wires.size() != 2)
return NaN;
184 double v = try_eval_constant(gc, wires[1]);
185 if (std::isnan(v) || v == 0.0)
return NaN;
186 return std::trunc(first / v);
189 if (wires.size() != 1)
return NaN;
193 for (std::size_t i = 1; i < wires.size(); ++i) {
194 double v = try_eval_constant(gc, wires[i]);
195 if (std::isnan(v))
return NaN;
202 for (std::size_t i = 1; i < wires.size(); ++i) {
203 double v = try_eval_constant(gc, wires[i]);
204 if (std::isnan(v))
return NaN;
210 if (wires.size() != 2)
return NaN;
211 double e = try_eval_constant(gc, wires[1]);
212 if (std::isnan(e))
return NaN;
218 return std::pow(first, e);
221 if (wires.size() != 1)
return NaN;
223 return std::log(first);
225 if (wires.size() != 1)
return NaN;
226 return std::exp(first);
247bool subtree_contains_agg(
const GenericCircuit &gc,
gate_t g)
249 std::unordered_set<gate_t> seen;
250 std::stack<gate_t> stk;
252 while (!stk.empty()) {
253 gate_t cur = stk.top(); stk.pop();
254 if (!seen.insert(cur).second)
continue;
268void replace_with_value(GenericCircuit &gc,
gate_t g,
double c)
277bool is_value_equal_to(
const GenericCircuit &gc,
gate_t g,
double target)
281 catch (
const CircuitException &) {
return false; }
300bool try_identity_drop(GenericCircuit &gc,
gate_t g)
306 std::vector<gate_t> kept;
307 kept.reserve(wires.size());
309 if (!is_value_equal_to(gc, w, 0.0)) kept.push_back(w);
311 if (kept.size() == wires.size())
return false;
313 replace_with_value(gc, g, 0.0);
316 wires = std::move(kept);
322 if (is_value_equal_to(gc, w, 0.0)) {
323 replace_with_value(gc, g, 0.0);
327 std::vector<gate_t> kept;
328 kept.reserve(wires.size());
330 if (!is_value_equal_to(gc, w, 1.0)) kept.push_back(w);
332 if (kept.size() == wires.size())
return false;
334 replace_with_value(gc, g, 1.0);
337 wires = std::move(kept);
363bool is_invalid(
gate_t g) {
return g == INVALID_GATE; }
390std::optional<LinearTerm>
391decompose_linear_term(
const GenericCircuit &gc,
gate_t g)
398 catch (
const CircuitException &) {
return std::nullopt; }
399 return LinearTerm{INVALID_GATE, 0.0, v};
406 return LinearTerm{g, 1.0, 0.0};
421 return LinearTerm{g, 1.0, 0.0};
437 && wires.size() == 1) {
438 return decompose_linear_term(gc, wires[0]);
442 if (wires.size() != 1)
return std::nullopt;
443 auto inner = decompose_linear_term(gc, wires[0]);
444 if (!inner)
return std::nullopt;
445 return LinearTerm{inner->rv_gate, -inner->a, -inner->b};
449 if (wires.size() != 2)
return std::nullopt;
452 gate_t var_side = INVALID_GATE;
455 catch (
const CircuitException &) {
return std::nullopt; }
459 catch (
const CircuitException &) {
return std::nullopt; }
464 auto inner = decompose_linear_term(gc, var_side);
465 if (!inner)
return std::nullopt;
466 return LinearTerm{inner->rv_gate, c * inner->a, c * inner->b};
506bool try_sum_closure(GenericCircuit &gc,
gate_t g)
509 if (wires.size() < 2)
return false;
511 std::vector<LinearTerm> lterms;
512 lterms.reserve(wires.size());
514 auto term = decompose_linear_term(gc, w);
515 if (!term)
return false;
516 lterms.push_back(*term);
520 std::vector<std::unique_ptr<Distribution>> dists(lterms.size());
521 std::vector<ClosureTerm> terms;
522 terms.reserve(lterms.size());
523 std::unordered_set<gate_t> seen_rvs;
524 for (std::size_t i = 0; i < lterms.size(); ++i) {
525 const auto &t = lterms[i];
526 if (is_invalid(t.rv_gate)) {
527 terms.push_back({
nullptr, t.a, t.b});
530 if (!seen_rvs.insert(t.rv_gate).second)
return false;
534 if (is_shared_base_rv(t.rv_gate))
return false;
536 if (!spec)
return false;
538 terms.push_back({dists[i].get(), t.a, t.b});
542 if (!folded)
return false;
560bool try_product_closure(GenericCircuit &gc,
gate_t g)
563 if (wires.size() < 2)
return false;
565 double c_total = 1.0;
566 std::vector<std::unique_ptr<Distribution>> dists;
567 std::vector<const Distribution *> factors;
568 std::unordered_set<gate_t> seen_rvs;
573 catch (
const CircuitException &) {
return false; }
576 if (t !=
gate_rv)
return false;
577 if (!seen_rvs.insert(w).second)
return false;
578 if (is_shared_base_rv(w))
return false;
580 if (!spec)
return false;
582 factors.push_back(dists.back().get());
584 if (factors.size() < 2)
return false;
587 if (!combined)
return false;
588 if (c_total != 1.0) {
589 combined = combined->scale(c_total);
590 if (!combined)
return false;
607bool try_transform_closure(GenericCircuit &gc,
gate_t g)
613 if (!transform)
return false;
615 if (wires.size() != 1)
return false;
617 if (is_shared_base_rv(wires[0]))
return false;
619 if (!spec)
return false;
622 if (!image)
return false;
643bool try_neg_rv(GenericCircuit &gc,
gate_t g)
649 if (wires.size() != 1)
return false;
651 if (is_shared_base_rv(wires[0]))
return false;
654 if (!spec)
return false;
657 if (!negated)
return false;
693unsigned apply_rules(GenericCircuit &gc,
gate_t g,
694 bool include_scalar_fold);
713bool try_categorical_mixture_lift(GenericCircuit &gc,
gate_t g,
716 const std::vector<gate_t> &others)
729 catch (
const CircuitException &) {
return false; }
743 const std::vector<gate_t> mw = gc.
getWires(mix_gate);
745 std::vector<gate_t> new_wires;
746 new_wires.reserve(mw.size());
747 new_wires.push_back(key);
748 for (std::size_t i = 1; i < mw.size(); ++i) {
749 const gate_t old_mul = mw[i];
752 catch (
const CircuitException &) {
return false; }
756 const double p = gc.
getProb(old_mul);
757 const auto vi =
static_cast<unsigned>(gc.
getInfos(old_mul).first);
760 new_wires.push_back(new_mul);
766bool try_mixture_lift(GenericCircuit &gc,
gate_t g,
767 bool include_scalar_fold)
773 if (wires.size() < 2)
return false;
776 std::size_t mix_idx =
static_cast<std::size_t
>(-1);
777 for (std::size_t i = 0; i < wires.size(); ++i) {
779 if (mix_idx !=
static_cast<std::size_t
>(-1))
return false;
783 if (mix_idx ==
static_cast<std::size_t
>(-1))
return false;
785 const auto mix_gate = wires[mix_idx];
790 std::vector<gate_t> others;
791 others.reserve(wires.size() - 1);
792 for (std::size_t i = 0; i < wires.size(); ++i) {
793 if (i != mix_idx) others.push_back(wires[i]);
800 return try_categorical_mixture_lift(gc, g, op, mix_gate, others);
804 const auto &mw = gc.
getWires(mix_gate);
805 if (mw.size() != 3)
return false;
806 const gate_t p_tok = mw[0];
807 const gate_t x_tok = mw[1];
808 const gate_t y_tok = mw[2];
814 std::vector<gate_t> new_x_wires = others; new_x_wires.push_back(x_tok);
815 std::vector<gate_t> new_y_wires = others; new_y_wires.push_back(y_tok);
831 apply_rules(gc, new_x, include_scalar_fold);
832 apply_rules(gc, new_y, include_scalar_fold);
862bool try_times_scalar_rv(GenericCircuit &gc,
gate_t g)
867 if (wires.size() != 2)
return false;
871 gate_t rv_side = INVALID_GATE;
875 catch (
const CircuitException &) {
return false; }
880 catch (
const CircuitException &) {
return false; }
888 if (c == 0.0 || c == 1.0)
return false;
890 if (is_shared_base_rv(rv_side))
return false;
893 if (!spec)
return false;
896 if (!scaled)
return false;
902 if (
auto dirac = scaled->asDirac()) {
903 replace_with_value(gc, g, *dirac);
940bool try_plus_aggregate(GenericCircuit &gc,
gate_t g,
941 bool include_scalar_fold)
945 const auto &wires_in = gc.
getWires(g);
946 if (wires_in.size() < 2)
return false;
948 std::vector<LinearTerm> terms;
949 terms.reserve(wires_in.size());
950 for (
gate_t w : wires_in) {
951 auto t = decompose_linear_term(gc, w);
952 if (!t)
return false;
959 std::vector<std::pair<gate_t, double>> coeffs;
960 double b_total = 0.0;
961 unsigned constants_in = 0;
962 for (
const auto &t : terms) {
964 if (is_invalid(t.rv_gate)) {
969 for (
auto &p : coeffs) {
970 if (p.first == t.rv_gate) {
976 if (!found) coeffs.emplace_back(t.rv_gate, t.a);
983 const bool has_duplicate = (coeffs.size() < terms.size() - constants_in);
984 const bool many_constants = (constants_in >= 2);
985 if (!has_duplicate && !many_constants)
return false;
988 std::vector<std::pair<gate_t, double>> kept;
989 kept.reserve(coeffs.size());
990 for (
const auto &p : coeffs) {
991 if (p.second != 0.0) kept.push_back(p);
996 replace_with_value(gc, g, b_total);
1014 if (kept.size() == 1 && b_total == 0.0) {
1015 const auto &only = kept.front();
1016 if (only.second == 1.0) {
1028 std::vector<gate_t> new_wires;
1029 new_wires.reserve(kept.size() + 1);
1030 for (
const auto &p : kept) {
1031 if (p.second == 1.0) {
1032 new_wires.push_back(p.first);
1037 new_wires.push_back(tm);
1040 if (b_total != 0.0) {
1044 gc.
setWires(g, std::move(new_wires));
1051 apply_rules(gc, w, include_scalar_fold);
1066unsigned apply_rules(GenericCircuit &gc,
gate_t g,
1067 bool include_scalar_fold)
1074 for (
unsigned iter = 0; iter < 32; ++iter) {
1079 double c = try_eval_constant(gc, g);
1080 if (!std::isnan(c)) {
1081 replace_with_value(gc, g, c);
1101 const auto &wires_in = gc.
getWires(g);
1102 if (wires_in.size() == 2) {
1103 const gate_t a = wires_in[0];
1104 const gate_t b = wires_in[1];
1130 const auto &wires_in = gc.
getWires(g);
1131 if (wires_in.size() == 2 && !subtree_contains_agg(gc, wires_in[0])) {
1132 const double c = try_eval_constant(gc, wires_in[1]);
1133 if (!std::isnan(c) && c != 0.0) {
1134 const gate_t x = wires_in[0];
1147 if (try_identity_drop(gc, g)) {
1158 if (try_mixture_lift(gc, g, include_scalar_fold)) {
1171 if (try_plus_aggregate(gc, g, include_scalar_fold)) {
1188 if (try_times_scalar_rv(gc, g)) {
1201 if (try_sum_closure(gc, g)) { ++local;
break; }
1204 if (try_product_closure(gc, g)) { ++local;
break; }
1207 if (try_transform_closure(gc, g)) { ++local;
break; }
1224void simplify(GenericCircuit &gc,
gate_t g,
1225 std::unordered_set<gate_t> &done,
unsigned &counter,
1226 bool include_scalar_fold)
1232 std::stack<std::pair<gate_t, std::size_t>> stk;
1233 if (!done.insert(g).second)
return;
1236 while (!stk.empty()) {
1237 auto &frame = stk.top();
1238 gate_t cur = frame.first;
1239 const auto &wires = gc.
getWires(cur);
1240 if (frame.second < wires.size()) {
1241 gate_t child = wires[frame.second++];
1242 if (done.insert(child).second) stk.emplace(child, 0);
1247 counter += apply_rules(gc, cur, include_scalar_fold);
1256 unsigned counter = 0;
1264 for (std::size_t i = 0; i < nb; ++i) {
1265 auto g =
static_cast<gate_t>(i);
1267 double c = try_eval_constant(gc, g);
1268 if (!std::isnan(c)) {
1269 replace_with_value(gc, g, c);
1278 unsigned counter = 0;
1280 for (std::size_t i = 0; i < nb; ++i) {
1281 auto g =
static_cast<gate_t>(i);
1287 const auto &wires = gc.
getWires(g);
1288 if (wires.size() != 3)
continue;
1301 const gate_t sel = wires[0];
1323 }
else if (pi == 0.0) {
1333 unsigned counter = 0;
1354 std::unordered_set<gate_t> shared_base_rvs;
1358 std::unordered_map<gate_t, unsigned> cmp_footprint_count;
1359 for (std::size_t i = 0; i < nb; ++i) {
1360 auto g =
static_cast<gate_t>(i);
1363 std::unordered_set<gate_t> fp;
1364 collect_reachable_base_rvs(gc, w, fp);
1365 for (
gate_t rv : fp) ++cmp_footprint_count[rv];
1368 for (
const auto &[rv, n] : cmp_footprint_count)
1369 if (n > 1) shared_base_rvs.insert(rv);
1371 for (std::size_t i = 0; i < nb; ++i) {
1372 auto g =
static_cast<gate_t>(i);
1377 if (arms.size() < 2)
continue;
1378 std::unordered_map<gate_t, unsigned> arm_count;
1379 for (
gate_t arm : arms) {
1380 std::unordered_set<gate_t> fp;
1381 collect_reachable_base_rvs(gc, arm, fp);
1382 for (
gate_t rv : fp) ++arm_count[rv];
1384 for (
const auto &[rv, n] : arm_count)
1385 if (n > 1) shared_base_rvs.insert(rv);
1388 g_shared_base_rvs = &shared_base_rvs;
1389 struct SharedGuard {
1390 ~SharedGuard() { g_shared_base_rvs =
nullptr; }
1403 std::unordered_set<gate_t> done;
1405 for (std::size_t i = 0; i < nb; ++i) {
1406 simplify(gc,
static_cast<gate_t>(i), done, counter,
1425 for (std::size_t i = 0; i < nb; ++i) {
1426 auto g =
static_cast<gate_t>(i);
1428 if (try_times_scalar_rv(gc, g)) ++counter;
1429 else if (try_neg_rv(gc, g)) ++counter;
1450 const auto &wires = gc.
getWires(cmp_gate);
1451 if (wires.size() != 2)
return false;
1453 std::unordered_set<gate_t> seen;
1454 std::stack<gate_t> stk;
1457 while (!stk.empty()) {
1458 gate_t g = stk.top(); stk.pop();
1459 if (!seen.insert(g).second)
continue;
1478 if (mw.size() != 3)
return false;
1496void collect_cmp_rv_footprint(
const GenericCircuit &gc,
gate_t cmp_gate,
1497 std::unordered_set<gate_t> &fp)
1499 std::unordered_set<gate_t> seen;
1500 std::stack<gate_t> stk;
1502 while (!stk.empty()) {
1503 gate_t g = stk.top(); stk.pop();
1504 if (!seen.insert(g).second)
continue;
1534 if (mw.size() == 3) { stk.push(mw[1]); stk.push(mw[2]); }
1563void collect_cmp_mixture_selectors(
const GenericCircuit &gc,
gate_t cmp_gate,
1564 std::unordered_set<gate_t> &sels)
1566 std::unordered_set<gate_t> seen;
1567 std::stack<gate_t> stk;
1569 while (!stk.empty()) {
1570 gate_t g = stk.top(); stk.pop();
1571 if (!seen.insert(g).second)
continue;
1580 if (mw.size() == 3) {
1602constexpr std::size_t JOINT_TABLE_K_MAX = 8;
1616bool is_analytic_singleton_cmp(
const GenericCircuit &gc,
gate_t cmp_gate)
1618 const auto &wires = gc.
getWires(cmp_gate);
1619 if (wires.size() != 2)
return false;
1660struct FastPathInfo {
1662 std::vector<ComparisonOperator> ops;
1663 std::vector<double> thresholds;
1709std::optional<FastPathInfo>
1710detect_shared_scalar(
const GenericCircuit &gc,
1711 const std::vector<gate_t> &cmps)
1714 info.ops.reserve(cmps.size());
1715 info.thresholds.reserve(cmps.size());
1719 const auto &wires = gc.
getWires(c);
1720 if (wires.size() != 2)
return std::nullopt;
1724 if (!ok)
return std::nullopt;
1729 return std::nullopt;
1732 double threshold = std::numeric_limits<double>::quiet_NaN();
1735 scalar_side = wires[0];
1737 catch (
const CircuitException &) {
return std::nullopt; }
1739 scalar_side = wires[1];
1741 catch (
const CircuitException &) {
return std::nullopt; }
1742 effective_op = flip_cmp_op(op);
1744 return std::nullopt;
1748 info.scalar = scalar_side;
1750 }
else if (info.scalar != scalar_side) {
1751 return std::nullopt;
1753 info.ops.push_back(effective_op);
1754 info.thresholds.push_back(threshold);
1785bool inline_fast_path(GenericCircuit &gc,
1786 const std::vector<gate_t> &cmps,
1787 const FastPathInfo &info,
1794 std::vector<double> ts = info.thresholds;
1795 std::sort(ts.begin(), ts.end());
1796 ts.erase(std::unique(ts.begin(), ts.end()), ts.end());
1797 const std::size_t m = ts.size();
1798 const std::size_t nb_intervals = m + 1;
1812 std::vector<double> interval_probs(nb_intervals, 0.0);
1813 bool analytical =
false;
1817 std::vector<double> cdf_at_boundary(m);
1819 for (std::size_t i = 0; i < m; ++i) {
1820 cdf_at_boundary[i] =
cdfAt(*spec, ts[i]);
1821 if (std::isnan(cdf_at_boundary[i])) { all_ok =
false;
break; }
1824 interval_probs[0] = cdf_at_boundary[0];
1825 for (std::size_t i = 1; i < m; ++i)
1826 interval_probs[i] = cdf_at_boundary[i] - cdf_at_boundary[i - 1];
1827 interval_probs[m] = 1.0 - cdf_at_boundary[m - 1];
1838 if (!allow_mc)
return false;
1840 for (
double s : draws) {
1841 auto it = std::upper_bound(ts.begin(), ts.end(), s);
1842 std::size_t idx =
static_cast<std::size_t
>(it - ts.begin());
1843 ++interval_probs[idx];
1845 for (
auto &p : interval_probs) p /= samples;
1854 std::vector<unsigned long> outcome_word(nb_intervals, 0);
1855 for (std::size_t i = 0; i < nb_intervals; ++i) {
1857 if (i == 0) point = ts[0] - 1.0;
1858 else if (i == m) point = ts[m - 1] + 1.0;
1859 else point = 0.5 * (ts[i - 1] + ts[i]);
1860 unsigned long w = 0;
1861 for (std::size_t j = 0; j < info.thresholds.size(); ++j) {
1862 if (apply_cmp(point, info.ops[j], info.thresholds[j]))
1865 outcome_word[i] = w;
1871 std::vector<gate_t> mul_for_interval(nb_intervals,
1872 static_cast<gate_t>(-1));
1873 for (std::size_t i = 0; i < nb_intervals; ++i) {
1874 if (interval_probs[i] <= 0.0)
continue;
1875 mul_for_interval[i] =
1877 static_cast<unsigned>(i));
1882 for (std::size_t j = 0; j < cmps.size(); ++j) {
1883 std::vector<gate_t> plus_wires;
1884 plus_wires.reserve(nb_intervals);
1885 for (std::size_t i = 0; i < nb_intervals; ++i) {
1886 if (!(outcome_word[i] & (1ul << j)))
continue;
1887 gate_t mw = mul_for_interval[i];
1888 if (mw ==
static_cast<gate_t>(-1))
continue;
1889 plus_wires.push_back(mw);
1916void inline_joint_table(GenericCircuit &gc,
1917 const std::vector<gate_t> &cmps,
1920 const unsigned k =
static_cast<unsigned>(cmps.size());
1935 const std::size_t nb_outcomes = std::size_t{1} << k;
1936 std::vector<gate_t> mul_for_outcome(nb_outcomes,
1937 static_cast<gate_t>(-1));
1938 for (std::size_t w = 0; w < nb_outcomes; ++w) {
1939 if (probs[w] <= 0.0)
continue;
1940 mul_for_outcome[w] =
1942 static_cast<unsigned>(w));
1947 for (
unsigned i = 0; i < k; ++i) {
1948 std::vector<gate_t> plus_wires;
1949 plus_wires.reserve(nb_outcomes / 2);
1950 for (std::size_t w = 0; w < nb_outcomes; ++w) {
1951 if ((w & (std::size_t{1} << i)) == 0)
continue;
1952 gate_t m = mul_for_outcome[w];
1953 if (m ==
static_cast<gate_t>(-1))
continue;
1954 plus_wires.push_back(m);
1969struct PivotIslandInfo {
1970 DistributionSpec pivotSpec;
1973 DistributionSpec other;
1977 std::vector<Factor> factors;
1984std::optional<PivotIslandInfo>
1985detect_shared_pivot_rv(
const GenericCircuit &gc,
1986 const std::vector<gate_t> &cmps)
1989 bool havePivot =
false;
1990 std::optional<DistributionSpec> pivotSpec;
1991 PivotIslandInfo info;
1992 std::unordered_set<gate_t> othersSeen;
1997 if (w.size() != 2)
return std::nullopt;
2001 return std::nullopt;
2004 gate_t pv, other;
bool pivotLeft;
2006 (!havePivot || w[0] == pivot)) { pv = w[0]; other = w[1]; pivotLeft =
true; }
2008 (!havePivot || w[1] == pivot)) { pv = w[1]; other = w[0]; pivotLeft =
false; }
2009 else return std::nullopt;
2013 if (!sp)
return std::nullopt;
2014 pivot = pv; pivotSpec = *sp; havePivot =
true;
2021 const bool trueIsGreater = pivotLeft ? greaterOp : lessOp;
2023 PivotIslandInfo::Factor f;
2024 f.trueIsGreater = trueIsGreater;
2028 catch (
const CircuitException &) {
return std::nullopt; }
2030 if (!othersSeen.insert(other).second)
return std::nullopt;
2032 if (!sp)
return std::nullopt;
2035 }
else return std::nullopt;
2036 info.factors.push_back(std::move(f));
2038 if (!havePivot)
return std::nullopt;
2039 info.pivotSpec = *pivotSpec;
2053bool inline_analytic_pivot_joint_table(GenericCircuit &gc,
2054 const std::vector<gate_t> &cmps,
2055 const PivotIslandInfo &info)
2057 const unsigned k =
static_cast<unsigned>(cmps.size());
2058 const std::size_t nb_outcomes = std::size_t{1} << k;
2062 if (!dX->integrationRange(lo0, hi0))
return false;
2065 std::vector<std::unique_ptr<Distribution>> otherDist(k);
2066 for (
unsigned j = 0; j < k; ++j)
2067 if (!info.factors[j].isConst)
2070 std::vector<double> probs(nb_outcomes, 0.0);
2071 for (std::size_t w = 0; w < nb_outcomes; ++w) {
2073 double lo = lo0, hi = hi0;
2074 for (
unsigned j = 0; j < k; ++j) {
2075 if (!info.factors[j].isConst)
continue;
2076 const bool bit = (w >> j) & 1u;
2078 const bool greater = (info.factors[j].trueIsGreater == bit);
2079 if (greater) lo = std::max(lo, info.factors[j].konst);
2080 else hi = std::min(hi, info.factors[j].konst);
2082 if (!(hi > lo)) { probs[w] = 0.0;
continue; }
2086 const double fX = dX->pdf(x);
2087 if (std::isnan(fX))
return std::numeric_limits<double>::quiet_NaN();
2089 for (
unsigned j = 0; j < k; ++j) {
2090 if (info.factors[j].isConst)
continue;
2091 const bool bit = (w >> j) & 1u;
2092 const bool greater = (info.factors[j].trueIsGreater == bit);
2093 const double FY = otherDist[j]->cdf(x);
2094 if (std::isnan(FY))
return std::numeric_limits<double>::quiet_NaN();
2095 weight *= greater ? FY : (1.0 - FY);
2099 if (std::isnan(cell))
return false;
2106 std::vector<gate_t> mul_for_outcome(nb_outcomes,
static_cast<gate_t>(-1));
2107 for (std::size_t w = 0; w < nb_outcomes; ++w) {
2108 if (probs[w] <= 0.0)
continue;
2109 mul_for_outcome[w] =
2112 for (
unsigned i = 0; i < k; ++i) {
2113 std::vector<gate_t> plus_wires;
2114 for (std::size_t w = 0; w < nb_outcomes; ++w) {
2115 if ((w & (std::size_t{1} << i)) == 0)
continue;
2116 gate_t mw = mul_for_outcome[w];
2117 if (mw ==
static_cast<gate_t>(-1))
continue;
2118 plus_wires.push_back(mw);
2140 const bool allow_mc = (samples > 0);
2149 std::vector<gate_t> cmps;
2150 for (std::size_t i = 0; i < nb; ++i) {
2151 auto g =
static_cast<gate_t>(i);
2158 std::unordered_map<gate_t, std::unordered_set<gate_t>> footprints;
2159 footprints.reserve(cmps.size());
2161 collect_cmp_rv_footprint(gc, c, footprints[c]);
2169 std::unordered_map<gate_t, unsigned> indeg;
2171 for (std::size_t i = 0; i < nb; ++i)
2180 std::vector<std::size_t> parent(cmps.size());
2181 for (std::size_t i = 0; i < cmps.size(); ++i) parent[i] = i;
2182 auto find = [&](std::size_t x) {
2183 while (parent[x] != x) {
2184 parent[x] = parent[parent[x]];
2189 auto unite = [&](std::size_t a, std::size_t b) {
2190 a = find(a); b = find(b);
2191 if (a != b) parent[a] = b;
2193 for (std::size_t i = 0; i < cmps.size(); ++i) {
2194 for (std::size_t j = i + 1; j < cmps.size(); ++j) {
2195 if (find(i) == find(j))
continue;
2196 const auto &fp_i = footprints[cmps[i]];
2197 const auto &fp_j = footprints[cmps[j]];
2198 const auto &small = fp_i.size() < fp_j.size() ? fp_i : fp_j;
2199 const auto &big = fp_i.size() < fp_j.size() ? fp_j : fp_i;
2200 for (
gate_t rv : small) {
2201 if (big.count(rv)) { unite(i, j);
break; }
2207 std::unordered_map<std::size_t, std::vector<gate_t>> groups;
2208 for (std::size_t i = 0; i < cmps.size(); ++i)
2209 groups[find(i)].push_back(cmps[i]);
2211 unsigned resolved = 0;
2212 for (
auto &[root, group] : groups) {
2217 bool all_pristine =
true;
2221 if (!all_pristine)
continue;
2240 std::unordered_set<gate_t> island;
2241 std::unordered_map<gate_t, unsigned> island_ref;
2243 std::stack<gate_t> stk;
2244 for (
gate_t c : group) stk.push(c);
2245 while (!stk.empty()) {
2246 gate_t g = stk.top(); stk.pop();
2247 if (!island.insert(g).second)
continue;
2248 for (
gate_t w : gc.
getWires(g)) { ++island_ref[w]; stk.push(w); }
2254 auto sub_dag_escapes_island = [&](
gate_t s) {
2255 std::unordered_set<gate_t> seen;
2256 std::stack<gate_t> st; st.push(s);
2257 while (!st.empty()) {
2258 gate_t g = st.top(); st.pop();
2259 if (!seen.insert(g).second)
continue;
2260 unsigned inside = island_ref.count(g) ? island_ref[g] : 0;
2261 unsigned total = indeg.count(g) ? indeg[g] : 0;
2262 if (total > inside)
return true;
2267 std::unordered_set<gate_t> sels;
2268 for (
gate_t c : group) collect_cmp_mixture_selectors(gc, c, sels);
2269 bool group_couples_selector =
false;
2271 if (sub_dag_escapes_island(s)) { group_couples_selector =
true;
break; }
2272 if (group_couples_selector)
continue;
2275 if (group.size() == 1) {
2282 if (is_analytic_singleton_cmp(gc, group[0]))
continue;
2286 if (!allow_mc)
continue;
2303 if (
auto info = detect_shared_scalar(gc, group)) {
2304 if (inline_fast_path(gc, group, *info, samples, allow_mc)) {
2305 resolved +=
static_cast<unsigned>(group.size());
2311 "the joint probability of correlated comparison events over a "
2312 "composite quantity needs Monte Carlo, but provsql.rv_mc_samples "
2313 "= 0 disables it; set provsql.rv_mc_samples > 0 (comparisons "
2314 "against constants on a single distribution stay analytical)");
2323 if (
auto pinfo = detect_shared_pivot_rv(gc, group)) {
2324 if (inline_analytic_pivot_joint_table(gc, group, *pinfo)) {
2325 resolved +=
static_cast<unsigned>(group.size());
2338 "the joint probability of correlated comparison events needs "
2339 "Monte Carlo, but provsql.rv_mc_samples = 0 disables it; set "
2340 "provsql.rv_mc_samples > 0 (comparisons against constants on a "
2341 "single distribution stay analytical)");
2343 if (group.size() > JOINT_TABLE_K_MAX)
continue;
2345 inline_joint_table(gc, group, samples);
2346 resolved +=
static_cast<unsigned>(group.size());
ComparisonOperator cmpOpFromOid(Oid op_oid, bool &ok)
Map a PostgreSQL comparison-operator OID to a ComparisonOperator.
Typed aggregation value, operator, and aggregator abstractions.
ComparisonOperator
SQL comparison operators used in gate_cmp circuit gates.
@ LE
Less than or equal (<=).
@ GE
Greater than or equal (>=).
Closed-form CDF resolution for trivial gate_cmp shapes.
gate_t
Strongly-typed gate identifier.
Per-family polymorphic view over a continuous gate_rv distribution (§F.1 class hierarchy).
Analytical expectation / variance / moment evaluator over RV circuits.
Peephole simplifier for continuous gate_arith sub-circuits.
Monte Carlo sampling over a GenericCircuit, RV-aware.
Shared 1-D quadrature core for the pivot-conjunction and order-statistic closed forms.
Continuous random-variable helpers (distribution parsing, moments).
Exception type thrown by circuit operations on invalid input.
std::vector< gate_t > & getWires(gate_t g)
Return a mutable reference to the child-wire list of gate g.
gateType getGateType(gate_t g) const
Return the type of gate g.
std::vector< gate_t >::size_type getNbGates() const
Return the total number of gates in the circuit.
In-memory provenance circuit with semiring-generic evaluation.
void resolveToPlus(gate_t g, std::vector< gate_t > w)
Rewrite an arbitrary gate as a gate_plus over w.
void resolveToCategoricalMixture(gate_t g, std::vector< gate_t > wires_)
Rewrite g in place as a categorical-form gate_mixture over wires ([key, mul_1, ......
void setWires(gate_t g, std::vector< gate_t > w)
Replace the wires of g with w.
gate_t addAnonymousMulinputGateWithValue(gate_t key, double p, unsigned value_index, const std::string &value_text)
Allocate a fresh gate_mulinput labelled with a numeric outcome value carried in extra.
void resolveToRv(gate_t g, const std::string &s)
Rewrite an arbitrary gate as a gate_rv carrying the distribution-spec extra s.
void resolveToMixture(gate_t g, gate_t p_token, gate_t x_token, gate_t y_token)
Rewrite g in place as a gate_mixture over the wires [p_token, x_token, y_token].
gate_t addAnonymousArithGate(provsql_arith_op op, std::vector< gate_t > wires_)
Allocate a fresh gate_arith gate with operator tag op and the given wires.
gate_t addAnonymousValueGate(const std::string &text)
Allocate a fresh gate_value gate carrying the textual scalar text.
bool isCategoricalMixture(gate_t g) const
Test whether g is a categorical-form gate_mixture (the explicit provsql.categorical output).
void setInfos(gate_t g, unsigned info1, unsigned info2)
Set the integer annotation pair for gate g.
std::string getExtra(gate_t g) const
Return the string extra for gate g.
double getProb(gate_t g) const
Return the probability for gate g.
void resolveCmpToBernoulli(gate_t g, double p)
Replace a gate_cmp by a constant Boolean leaf (gate_one for p == 1, gate_zero for p == 0) or by a Ber...
gate_t addAnonymousInputGate(double p)
Allocate a fresh gate_input gate carrying probability p, with a unique synthetic UUID so subsequent B...
std::pair< unsigned, unsigned > getInfos(gate_t g) const
Return the integer annotation pair for gate g.
void liftConditionedToTarget(gate_t g, gate_t target)
Replace a gate_conditioned g by a transparent passthrough to its target child (a single-child gate_ar...
gate_t addAnonymousMulinputGate(gate_t key, double p, unsigned value_index)
Allocate a fresh gate_mulinput gate with key key, probability p, and value index value_index.
void resolveToValue(gate_t g, const std::string &s)
Rewrite an arbitrary gate as a gate_value carrying the textual extra s.
unsigned runConstantFold(GenericCircuit &gc)
Constant-fold pass over every gate_arith in gc.
std::unique_ptr< Distribution > closeTransform(const char *transform, const Distribution &x)
The image distribution of transform applied to x, when a registered rule covers x's family; nullptr o...
double parseDoubleStrict(const std::string &s)
Strictly parse s as a double.
std::vector< double > monteCarloJointDistribution(const GenericCircuit &gc, const std::vector< gate_t > &cmps, unsigned samples)
Estimate the joint distribution of cmps via Monte Carlo.
std::unique_ptr< Distribution > makeDistribution(const DistributionSpec &spec)
Construct the per-family Distribution for a parsed spec.
std::unique_ptr< Distribution > closePlusTerms(const std::vector< ClosureTerm > &terms)
Fold PLUS(terms) into a single distribution when a registered closure covers every family in the sum.
double simpsonIntegrate(double lo, double hi, int N, F &&f)
Composite-Simpson with N panels.
unsigned runHybridSimplifier(GenericCircuit &gc)
Run the peephole simplifier over gc.
std::unique_ptr< Distribution > closeProductFactors(const std::vector< const Distribution * > &factors)
Fold a product of independent factors into a single distribution when a registered closure covers eve...
std::vector< double > monteCarloScalarSamples(const GenericCircuit &gc, gate_t root, unsigned samples)
Sample a scalar sub-circuit samples times and return the draws.
std::optional< DistributionSpec > parse_distribution_spec(const std::string &s)
Parse the on-disk text encoding of a gate_rv distribution.
double monteCarloRV(const GenericCircuit &gc, gate_t root, unsigned samples)
Run Monte Carlo on a circuit that may contain gate_rv leaves.
constexpr int kSimpsonPanels
Panel count shared by every composite-Simpson quadrature over a distribution's integration range: exa...
unsigned foldDegenerateMixtures(GenericCircuit &gc)
Collapse degenerate Bernoulli gate_mixture gates whose selector is certainly true (pi = 1) or certain...
double cdfAt(const DistributionSpec &d, double c)
Closed-form CDF for a basic continuous distribution.
std::string double_to_text(double v)
Format a double back into the canonical text form used by gate_value extras and gate_rv distribution ...
unsigned runHybridDecomposer(GenericCircuit &gc, unsigned samples)
Marginalise unresolved continuous-island gate_cmp gates into Bernoulli gate_input leaves.
Core types, constants, and utilities shared across ProvSQL.
provsql_arith_op
Arithmetic operator tags used by gate_arith.
@ PROVSQL_ARITH_PERCENTILE
continuous percentile (order-statistic aggregate): wires are interleaved [ind_1, x_1,...
@ PROVSQL_ARITH_DIV
binary, child0 / child1
@ PROVSQL_ARITH_ASFLOAT8
unary, child0 as double precision
@ PROVSQL_ARITH_LN
unary, natural logarithm of child0 (a negative draw raises at evaluation)
@ PROVSQL_ARITH_ROUND
child0 rounded half away from zero, to child1 decimal digits where a second child is given (SQL round...
@ PROVSQL_ARITH_PLUS
n-ary, sum of children
@ PROVSQL_ARITH_POW
binary, child0 ^ child1 (real branch only: a negative base drawn with a non-integer exponent raises a...
@ PROVSQL_ARITH_ABS
unary, |child0|
@ PROVSQL_ARITH_FLOOR
unary, greatest integer <= child0
@ PROVSQL_ARITH_NEG
unary, -child0
@ PROVSQL_ARITH_INTDIV
binary, child0 / child1 truncated toward zero: SQL's division of two integers
@ PROVSQL_ARITH_MINUS
binary, child0 - child1
@ PROVSQL_ARITH_CEIL
unary, least integer >= child0
@ PROVSQL_ARITH_EXP
unary, e^child0
@ PROVSQL_ARITH_TIMES
n-ary, product of children
@ PROVSQL_ARITH_MIN
n-ary, min of children (order statistic; least / min aggregate)
@ PROVSQL_ARITH_MAX
n-ary, max of children (order statistic; greatest / max aggregate)
@ PROVSQL_ARITH_ASFLOAT4
unary, child0 as real reads it
@ gate_rv
Continuous random-variable leaf (extra encodes distribution).
@ gate_mixture
Probabilistic mixture: three wires [p_token (gate_input Bernoulli), x_token, y_token]; samples x when...
@ gate_arith
n-ary arithmetic gate over scalar-valued children (info1 holds operator tag)