12#include <unordered_map>
13#include <unordered_set>
31#include "access/htup_details.h"
32#include "utils/uuid.h"
55 static Interval point(
double v) {
return {v, v}; }
56 static Interval all() {
return {-std::numeric_limits<double>::infinity(),
57 +std::numeric_limits<double>::infinity()}; }
59 return std::isinf(lo) && lo < 0 && std::isinf(hi) && hi > 0;
63Interval add(Interval a, Interval b) {
return {a.lo + b.lo, a.hi + b.hi}; }
64Interval sub(Interval a, Interval b) {
return {a.lo - b.hi, a.hi - b.lo}; }
65Interval neg(Interval a) {
return {-a.hi, -a.lo}; }
69Interval mul(Interval a, Interval b)
71 double p1 = a.lo * b.lo, p2 = a.lo * b.hi;
72 double p3 = a.hi * b.lo, p4 = a.hi * b.hi;
73 return {std::min({p1, p2, p3, p4}), std::max({p1, p2, p3, p4})};
82Interval expInt(Interval a) {
return {std::exp(a.lo), std::exp(a.hi)}; }
87Interval lnInt(Interval a)
89 const double lo = a.lo > 0.0 ? std::log(a.lo)
90 : -std::numeric_limits<double>::infinity();
91 const double hi = a.hi > 0.0 ? std::log(a.hi)
92 : -std::numeric_limits<double>::infinity();
102Interval powInt(Interval b, Interval e)
105 return Interval::all();
106 double lo = std::numeric_limits<double>::infinity();
107 double hi = -std::numeric_limits<double>::infinity();
108 for (
double x : {b.lo, b.hi})
109 for (
double y : {e.lo, e.hi}) {
110 const double v = std::pow(x, y);
112 return Interval::all();
113 lo = std::min(lo, v);
114 hi = std::max(hi, v);
119Interval divInt(Interval a, Interval b)
121 if (b.lo <= 0.0 && b.hi >= 0.0)
122 return Interval::all();
123 Interval inv = {1.0 / b.hi, 1.0 / b.lo};
141Interval intervalOf(
const GenericCircuit &gc,
gate_t g,
142 std::unordered_map<gate_t, Interval> &
cache)
144 auto it =
cache.find(g);
145 if (it !=
cache.
end())
return it->second;
147 Interval result = Interval::all();
158 }
catch (
const CircuitException &) {
159 result = Interval::all();
167 result = {s.lo, s.hi};
173 if (wires.empty())
break;
174 Interval first = intervalOf(gc, wires[0],
cache);
178 for (std::size_t i = 1; i < wires.size(); ++i)
179 result = add(result, intervalOf(gc, wires[i],
cache));
183 for (std::size_t i = 1; i < wires.size(); ++i)
184 result = mul(result, intervalOf(gc, wires[i],
cache));
187 if (wires.size() != 2)
break;
188 result = sub(first, intervalOf(gc, wires[1],
cache));
191 if (wires.size() != 2)
break;
192 result = divInt(first, intervalOf(gc, wires[1],
cache));
196 if (wires.size() != 2)
break;
197 result = divInt(first, intervalOf(gc, wires[1],
cache));
198 if (std::isfinite(result.lo)) result.lo = std::trunc(result.lo);
199 if (std::isfinite(result.hi)) result.hi = std::trunc(result.hi);
202 if (wires.size() != 1)
break;
210 for (std::size_t i = 1; i < wires.size(); ++i) {
211 Interval o = intervalOf(gc, wires[i],
cache);
212 result = { std::max(result.lo, o.lo), std::max(result.hi, o.hi) };
217 for (std::size_t i = 1; i < wires.size(); ++i) {
218 Interval o = intervalOf(gc, wires[i],
cache);
219 result = { std::min(result.lo, o.lo), std::min(result.hi, o.hi) };
223 if (wires.size() != 2)
break;
224 result = powInt(first, intervalOf(gc, wires[1],
cache));
227 if (wires.size() != 1)
break;
228 result = lnInt(first);
231 if (wires.size() != 1)
break;
232 result = expInt(first);
239 result = { std::floor(first.lo), std::ceil(first.hi) };
242 if (wires.size() != 1)
break;
243 result = { std::floor(first.lo), std::floor(first.hi) };
246 if (wires.size() != 1)
break;
247 result = { std::ceil(first.lo), std::ceil(first.hi) };
252 if (wires.size() != 1)
break;
256 if (wires.size() != 1)
break;
257 result = {
static_cast<double>(
static_cast<float>(first.lo)),
258 static_cast<double>(
static_cast<float>(first.hi)) };
263 if (wires.size() != 1)
break;
266 else if (first.hi <= 0.0)
267 result = { -first.hi, -first.lo };
269 result = { 0.0, std::max(-first.lo, first.hi) };
276 if (wires.size() < 2 || wires.size() % 2 != 0)
break;
277 result = intervalOf(gc, wires[1],
cache);
278 for (std::size_t i = 3; i < wires.size(); i += 2) {
279 Interval o = intervalOf(gc, wires[i],
cache);
280 result = { std::min(result.lo, o.lo), std::max(result.hi, o.hi) };
296 result = intervalOf(gc, wires[1],
cache);
310 double lo = std::numeric_limits<double>::infinity();
311 double hi = -std::numeric_limits<double>::infinity();
313 for (std::size_t i = 1; i < wires.size(); ++i) {
316 catch (
const CircuitException &) { any =
false;
break; }
317 lo = std::min(lo, v);
318 hi = std::max(hi, v);
321 if (any) result = {lo, hi};
322 }
else if (wires.size() == 3) {
323 Interval ix = intervalOf(gc, wires[1],
cache);
324 Interval iy = intervalOf(gc, wires[2],
cache);
325 result = {std::min(ix.lo, iy.lo), std::max(ix.hi, iy.hi)};
335 if (wires.empty())
break;
336 result = intervalOf(gc, wires.back(),
cache);
337 for (std::size_t i = 1; i + 1 < wires.size(); i += 2) {
338 Interval o = intervalOf(gc, wires[i],
cache);
339 result = { std::min(result.lo, o.lo), std::max(result.hi, o.hi) };
370 if (diff.hi < 0.0)
return 1.0;
371 if (diff.lo >= 0.0)
return 0.0;
374 if (diff.hi <= 0.0)
return 1.0;
375 if (diff.lo > 0.0)
return 0.0;
378 if (diff.lo > 0.0)
return 1.0;
379 if (diff.hi <= 0.0)
return 0.0;
382 if (diff.lo >= 0.0)
return 1.0;
383 if (diff.hi < 0.0)
return 0.0;
391 if (diff.hi < 0.0 || diff.lo > 0.0)
return 0.0;
394 if (diff.hi < 0.0 || diff.lo > 0.0)
return 1.0;
397 return std::numeric_limits<double>::quiet_NaN();
431double decideAggVsConstCmp(
const GenericCircuit &gc,
gate_t agg_gate,
434 bool *out_always_true =
nullptr)
439 std::vector<double> values;
442 return std::numeric_limits<double>::quiet_NaN();
443 const auto &sw = gc.
getWires(child);
445 return std::numeric_limits<double>::quiet_NaN();
448 return std::numeric_limits<double>::quiet_NaN();
451 }
catch (
const CircuitException &) {
452 return std::numeric_limits<double>::quiet_NaN();
456 Interval val_interval = Interval::all();
460 val_interval = {0.0,
static_cast<double>(values.size())};
463 double sum_neg = 0.0, sum_pos = 0.0;
464 for (
double v : values) {
465 if (v < 0.0) sum_neg += v;
468 val_interval = {std::min(0.0, sum_neg), std::max(0.0, sum_pos)};
474 return std::numeric_limits<double>::quiet_NaN();
475 val_interval = {*std::min_element(values.begin(), values.end()),
476 *std::max_element(values.begin(), values.end())};
481 return std::numeric_limits<double>::quiet_NaN();
484 Interval lhs = agg_on_lhs ? val_interval : Interval::point(const_val);
485 Interval rhs = agg_on_lhs ? Interval::point(const_val) : val_interval;
486 Interval diff = sub(lhs, rhs);
487 double p = decideCmp(diff, op);
498 if (p == 0.0)
return 0.0;
499 if (p == 1.0 && out_always_true !=
nullptr) *out_always_true =
true;
500 return std::numeric_limits<double>::quiet_NaN();
513double extractScalarConst(
const GenericCircuit &gc,
gate_t g)
518 catch (
const CircuitException &) {
519 return std::numeric_limits<double>::quiet_NaN();
524 if (w.size() != 2)
return std::numeric_limits<double>::quiet_NaN();
526 return std::numeric_limits<double>::quiet_NaN();
528 return std::numeric_limits<double>::quiet_NaN();
530 catch (
const CircuitException &) {
531 return std::numeric_limits<double>::quiet_NaN();
534 return std::numeric_limits<double>::quiet_NaN();
569bool asRvVsConstCmp(
const GenericCircuit &gc,
gate_t cmp_gate,
575 if (!ok)
return false;
576 const auto &wires = gc.
getWires(cmp_gate);
577 if (wires.size() != 2)
return false;
594 catch (
const CircuitException &) {
return false; }
601 catch (
const CircuitException &) {
return false; }
603 op_out = flipCmpOp(op);
625 current.hi = std::min(current.hi, c);
629 current.lo = std::max(current.lo, c);
632 current.lo = std::max(current.lo, c);
633 current.hi = std::min(current.hi, c);
643bool intervalEmpty(Interval i) {
return i.lo > i.hi; }
669void walkAndConjunctIntervals(
670 const GenericCircuit &gc,
gate_t root,
671 std::unordered_map<gate_t, Interval> &rv_intervals,
672 std::unordered_map<gate_t, Interval> &support_cache,
675 std::unordered_set<gate_t> seen;
676 std::stack<gate_t> stk;
680 while (!stk.empty()) {
681 gate_t g = stk.top(); stk.pop();
682 if (!seen.insert(g).second)
continue;
689 if (!asRvVsConstCmp(gc, g, rv, op, c)) {
697 auto it = rv_intervals.find(rv);
698 Interval current = (it == rv_intervals.end())
699 ? intervalOf(gc, rv, support_cache)
701 current = intersectRvConstraint(current, op, c);
702 rv_intervals[rv] = current;
752bool isAndJointlyInfeasible(
const GenericCircuit &gc,
gate_t root)
754 std::unordered_map<gate_t, Interval> rv_intervals;
755 std::unordered_map<gate_t, Interval> support_cache;
757 walkAndConjunctIntervals(gc, root, rv_intervals, support_cache, complete);
758 for (
const auto &kv : rv_intervals) {
759 if (intervalEmpty(kv.second))
return true;
802bool hasOnlyContinuousSupport(
const GenericCircuit &gc,
gate_t g,
803 std::unordered_map<gate_t, bool> &
cache)
805 auto it =
cache.find(g);
806 if (it !=
cache.
end())
return it->second;
822 result = !(tmpl && tmpl->family->factory(0.0, 0.0)->isDiscrete());
831 if (!hasOnlyContinuousSupport(gc, w,
cache)) { result =
false;
break; }
838 if (w.size() != 3) { result =
false;
break; }
839 result = hasOnlyContinuousSupport(gc, w[1],
cache)
840 && hasOnlyContinuousSupport(gc, w[2],
cache);
873const std::unordered_set<gate_t> &
874collectRandomLeaves(
const GenericCircuit &gc,
gate_t g,
875 std::unordered_map<
gate_t, std::unordered_set<gate_t>> &
cache)
877 auto it =
cache.find(g);
878 if (it !=
cache.
end())
return it->second;
885 cache.emplace(g, std::unordered_set<gate_t>{});
887 std::unordered_set<gate_t> out;
893 const auto &child = collectRandomLeaves(gc, w,
cache);
894 out.insert(child.begin(), child.end());
901 auto fit =
cache.find(g);
902 fit->second = std::move(out);
906using DiracMap = std::unordered_map<double, double>;
907using DiracMapOpt = std::optional<DiracMap>;
941collectDiracMassMap(
const GenericCircuit &gc,
gate_t g,
942 std::unordered_map<gate_t, DiracMapOpt> &
cache)
944 auto it =
cache.find(g);
945 if (it !=
cache.
end())
return it->second;
947 cache.emplace(g, std::nullopt);
956 result = std::move(m);
957 }
catch (
const CircuitException &) {
971 if (tmpl && tmpl->family->factory(0.0, 0.0)->isDiscrete())
972 result = std::nullopt;
982 for (std::size_t i = 1; i < w.size(); ++i) {
985 catch (
const CircuitException &) { ok =
false;
break; }
986 const double p = gc.
getProb(w[i]);
987 if (!std::isfinite(p) || p < 0.0 || p > 1.0) { ok =
false;
break; }
990 if (ok) result = std::move(m);
991 }
else if (w.size() == 3
993 const double pi = gc.
getProb(w[0]);
994 if (std::isfinite(pi) && pi >= 0.0 && pi <= 1.0) {
995 auto mx = collectDiracMassMap(gc, w[1],
cache);
996 auto my = collectDiracMassMap(gc, w[2],
cache);
999 for (
const auto &[v, mass] : *mx) m[v] += pi * mass;
1000 for (
const auto &[v, mass] : *my) m[v] += (1.0 - pi) * mass;
1001 result = std::move(m);
1011 auto fit =
cache.find(g);
1012 fit->second = result;
1020 std::unordered_map<gate_t, Interval>
cache;
1025 std::unordered_map<gate_t, bool> continuous_support_cache;
1026 std::unordered_map<gate_t, DiracMapOpt> dirac_cache;
1027 std::unordered_map<gate_t, std::unordered_set<gate_t>> leaf_cache;
1028 unsigned resolved = 0;
1036 std::vector<gate_t> cmps;
1037 for (std::size_t i = 0; i < nb; ++i) {
1038 auto g =
static_cast<gate_t>(i);
1050 const auto &wires = gc.
getWires(c);
1051 if (wires.size() != 2)
continue;
1060 if (wires[0] == wires[1]) {
1061 double p = std::numeric_limits<double>::quiet_NaN();
1098 bool lhs_continuous = hasOnlyContinuousSupport(gc, wires[0],
1099 continuous_support_cache);
1100 bool rhs_continuous = hasOnlyContinuousSupport(gc, wires[1],
1101 continuous_support_cache);
1102 if (lhs_continuous || rhs_continuous) {
1130 auto m_l = collectDiracMassMap(gc, wires[0], dirac_cache);
1131 auto m_r = collectDiracMassMap(gc, wires[1], dirac_cache);
1133 const auto &leaves_l = collectRandomLeaves(gc, wires[0], leaf_cache);
1134 const auto &leaves_r = collectRandomLeaves(gc, wires[1], leaf_cache);
1135 bool independent =
true;
1136 for (
gate_t leaf : leaves_l) {
1137 if (leaves_r.count(leaf)) { independent =
false;
break; }
1143 const DiracMap *small = (m_l->size() <= m_r->size()) ? &*m_l : &*m_r;
1144 const DiracMap *large = (m_l->size() <= m_r->size()) ? &*m_r : &*m_l;
1145 for (
const auto &[v, mass] : *small) {
1146 auto fit = large->find(v);
1147 if (fit != large->end()) p_eq += mass * fit->second;
1152 if (p_eq < 0.0) p_eq = 0.0;
1153 if (p_eq > 1.0) p_eq = 1.0;
1168 if (lhs_is_agg != rhs_is_agg) {
1169 gate_t agg_side = lhs_is_agg ? wires[0] : wires[1];
1170 gate_t const_side = lhs_is_agg ? wires[1] : wires[0];
1171 double const_val = extractScalarConst(gc, const_side);
1172 if (!std::isnan(const_val)) {
1173 double p = decideAggVsConstCmp(gc, agg_side, op, const_val,
1175 if (!std::isnan(p)) {
1184 Interval lhs = intervalOf(gc, wires[0],
cache);
1185 Interval rhs = intervalOf(gc, wires[1],
cache);
1188 if (lhs.isAll() && rhs.isAll())
continue;
1190 Interval diff = sub(lhs, rhs);
1191 double p = decideCmp(diff, op);
1192 if (!std::isnan(p)) {
1210 std::vector<gate_t> times_gates;
1211 for (std::size_t i = 0; i < nb_after; ++i) {
1212 auto g =
static_cast<gate_t>(i);
1214 times_gates.push_back(g);
1216 for (
gate_t t : times_gates) {
1218 if (isAndJointlyInfeasible(gc, t)) {
1282 unsigned resolved = 0;
1284 std::vector<gate_t> cmps;
1285 for (std::size_t i = 0; i < nb; ++i) {
1286 auto g =
static_cast<gate_t>(i);
1296 const auto &wires = gc.
getWires(c);
1297 if (wires.size() != 2)
continue;
1308 if (!(wires[0] == wires[1] && lhs_is_agg))
continue;
1314 std::vector<gate_t> ks;
1315 bool shape_ok =
true;
1327 ks.push_back(sw[0]);
1329 if (!shape_ok)
continue;
1331 if (!reflexive_true) {
1345 if (ks.empty())
continue;
1357 unsigned resolved = 0;
1360 std::vector<gate_t> cmps;
1361 cmps.reserve(nb / 8);
1362 for (std::size_t i = 0; i < nb; ++i) {
1363 auto g =
static_cast<gate_t>(i);
1375 const auto &wires = gc.
getWires(c);
1376 if (wires.size() != 2)
continue;
1381 if (lhs_is_agg == rhs_is_agg)
continue;
1383 gate_t agg_side = lhs_is_agg ? wires[0] : wires[1];
1384 gate_t const_side = lhs_is_agg ? wires[1] : wires[0];
1386 double const_val = extractScalarConst(gc, const_side);
1387 if (std::isnan(const_val))
continue;
1389 bool always_true =
false;
1390 double p = decideAggVsConstCmp(gc, agg_side, op, const_val,
1391 lhs_is_agg, &always_true);
1417 std::vector<gate_t> ks;
1418 bool shape_ok =
true;
1419 ks.reserve(gc.
getWires(agg_side).size());
1423 if (sw.size() != 2) { shape_ok =
false;
break; }
1424 ks.push_back(sw[0]);
1426 if (!shape_ok || ks.empty())
continue;
1453 std::unordered_map<gate_t, Interval> &
cache)
1457 const auto &wires = gc.
getWires(root);
1458 if (wires.size() != 2)
return Interval::all();
1459 Interval vi = intervalOf(gc, wires[1],
cache);
1463 return Interval{std::min(0.0, vi.lo), std::max(0.0, vi.hi)};
1467 const auto &wires = gc.
getWires(root);
1471 std::vector<gate_t> sm_children;
1472 sm_children.reserve(wires.size());
1476 auto child_value_iv = [&](
gate_t sm) -> Interval {
1478 if (sw.size() != 2)
return Interval::all();
1479 return intervalOf(gc, sw[1],
cache);
1481 auto child_always_fires = [&](
gate_t sm) ->
bool {
1486 const auto inf = std::numeric_limits<double>::infinity();
1490 return Interval{0.0,
static_cast<double>(sm_children.size())};
1496 double lo = 0.0, hi = 0.0;
1497 for (
gate_t sm : sm_children) {
1498 Interval vi = child_value_iv(sm);
1499 if (vi.isAll())
return Interval::all();
1500 if (child_always_fires(sm)) {
1504 lo += std::min(0.0, vi.lo);
1505 hi += std::max(0.0, vi.hi);
1508 return Interval{lo, hi};
1517 if (sm_children.empty())
return Interval::all();
1520 for (
gate_t sm : sm_children) {
1521 Interval vi = child_value_iv(sm);
1522 if (vi.isAll())
return Interval::all();
1523 lo = std::min(lo, vi.lo);
1524 hi = std::max(hi, vi.hi);
1526 if (lo > hi)
return Interval::all();
1527 return Interval{lo, hi};
1533 return Interval::all();
1539std::pair<double, double>
1541 std::optional<gate_t> event_root)
1543 std::unordered_map<gate_t, Interval>
cache;
1544 Interval iv = aggSupportOf(gc, root,
cache);
1555 if (event_root.has_value()) {
1556 std::unordered_map<gate_t, Interval> rv_intervals;
1558 walkAndConjunctIntervals(gc, *event_root, rv_intervals,
cache, complete);
1559 auto it = rv_intervals.find(root);
1560 if (it != rv_intervals.end()) {
1561 iv.lo = std::max(iv.lo, it->second.lo);
1562 iv.hi = std::min(iv.hi, it->second.hi);
1565 if (iv.lo > iv.hi) iv.lo = iv.hi;
1569 return {iv.lo, iv.hi};
1572std::optional<std::pair<double, double>>
1576 std::unordered_map<gate_t, Interval> rv_intervals;
1577 std::unordered_map<gate_t, Interval> support_cache;
1579 walkAndConjunctIntervals(gc, event_root, rv_intervals, support_cache,
1581 if (!complete)
return std::nullopt;
1588 auto it = rv_intervals.find(target_rv);
1590 if (it != rv_intervals.end()) {
1594 Interval base = intervalOf(gc, target_rv, support_cache);
1595 iv.lo = std::max(iv.lo, base.lo);
1596 iv.hi = std::min(iv.hi, base.hi);
1597 if (iv.lo > iv.hi) iv.lo = iv.hi;
1599 iv = intervalOf(gc, target_rv, support_cache);
1601 return std::make_pair(iv.lo, iv.hi);
1616 const std::string &s = gc.
getExtra(x);
1617 if (s.empty())
return false;
1620 double v = std::stod(s, &idx);
1621 if (idx != s.size() || !std::isfinite(v))
return false;
1634 const std::string &s = gc.
getExtra(mul);
1635 if (s.empty())
return false;
1638 double v = std::stod(s, &idx);
1639 if (idx != s.size() || !std::isfinite(v))
return false;
1647std::optional<TruncatedSingleRv>
1649 std::optional<gate_t> event_root)
1653 if (!spec)
return std::nullopt;
1662 double nat_lo = nat_support.
lo;
1663 double nat_hi = nat_support.
hi;
1666 if (!event_root.has_value()
1679 if (!iv.has_value())
return std::nullopt;
1680 if (!(iv->first < iv->second))
return std::nullopt;
1686 std::optional<gate_t> event_root)
1688 if (!event_root.has_value())
return false;
1702 if (!iv.has_value())
return false;
1703 return !(iv->first < iv->second);
1722static std::optional<double>
1725 return std::visit([&](
const auto &v) -> std::optional<double> {
1726 using T = std::decay_t<
decltype(v)>;
1727 if constexpr (std::is_same_v<T, TruncatedSingleRv>) {
1728 const double a = std::max(lo, v.lo);
1729 const double b = std::min(hi, v.hi);
1730 if (!(a < b))
return 0.0;
1731 const double cl = std::isfinite(a) ?
cdfAt(v.spec, a) : 0.0;
1732 const double ch = std::isfinite(b) ?
cdfAt(v.spec, b) : 1.0;
1733 if (std::isnan(cl) || std::isnan(ch))
return std::nullopt;
1735 }
else if constexpr (std::is_same_v<T, DiracShape>) {
1736 return (v.value >= lo && v.value <= hi) ? 1.0 : 0.0;
1737 }
else if constexpr (std::is_same_v<T, CategoricalShape>) {
1739 for (
const auto &pr : v.outcomes)
1740 if (pr.first >= lo && pr.first <= hi) m += pr.second;
1742 }
else if constexpr (std::is_same_v<T, BernoulliMixtureShape>) {
1745 if (!L || !R)
return std::nullopt;
1746 return v.p * (*L) + (1.0 - v.p) * (*R);
1748 return std::nullopt;
1768static std::optional<ClosedFormShape>
1771 return std::visit([&](
const auto &v) -> std::optional<ClosedFormShape> {
1772 using T = std::decay_t<
decltype(v)>;
1773 if constexpr (std::is_same_v<T, TruncatedSingleRv>) {
1774 const double a = std::max(lo, v.lo);
1775 const double b = std::min(hi, v.hi);
1776 if (!(a < b))
return std::nullopt;
1778 }
else if constexpr (std::is_same_v<T, DiracShape>) {
1779 if (v.value < lo || v.value > hi)
return std::nullopt;
1781 }
else if constexpr (std::is_same_v<T, CategoricalShape>) {
1784 for (
const auto &pr : v.outcomes) {
1785 if (pr.first >= lo && pr.first <= hi) {
1786 out.
outcomes.emplace_back(pr.first, pr.second);
1790 if (out.
outcomes.empty() || !(total > 0.0))
return std::nullopt;
1791 for (
auto &pr : out.
outcomes) pr.second /= total;
1793 }
else if constexpr (std::is_same_v<T, BernoulliMixtureShape>) {
1796 if (!mL || !mR)
return std::nullopt;
1797 const double pL = v.p * (*mL);
1798 const double pR = (1.0 - v.p) * (*mR);
1799 const double Z = pL + pR;
1800 if (!(Z > 0.0))
return std::nullopt;
1806 if (!Lt && !Rt)
return std::nullopt;
1811 m.
left = std::make_shared<ClosedFormShape>(std::move(*Lt));
1812 m.
right = std::make_shared<ClosedFormShape>(std::move(*Rt));
1815 return std::nullopt;
1819std::optional<ClosedFormShape>
1821 std::optional<gate_t> event_root)
1825 const bool event_trivial = !event_root.has_value()
1845 if (!m)
return std::nullopt;
1854 auto with_optional_truncation =
1855 [&](std::optional<ClosedFormShape> unc)
1856 -> std::optional<ClosedFormShape> {
1857 if (!unc)
return std::nullopt;
1858 if (event_trivial)
return unc;
1860 if (!iv.has_value())
return std::nullopt;
1861 if (!(iv->first < iv->second))
return std::nullopt;
1887 for (std::size_t i = 1; i < w.size(); ++i) {
1891 if (!std::isfinite(p) || p < 0.0 || p > 1.0)
return std::nullopt;
1894 if (cs.
outcomes.empty())
return std::nullopt;
1902 if (w.size() != 3)
return std::nullopt;
1905 if (!std::isfinite(p) || p < 0.0 || p > 1.0)
return std::nullopt;
1909 if (!left || !right)
return std::nullopt;
1913 m.
left = std::make_shared<ClosedFormShape>(std::move(*left));
1914 m.
right = std::make_shared<ClosedFormShape>(std::move(*right));
1918 return std::nullopt;
1944 gate_t root_gate, event_gate;
1951 std::optional<gate_t> event_opt;
1953 event_opt = event_gate;
1964 bool nulls[2] = {
false,
false};
1966 if (get_call_result_type(fcinfo, NULL, &tupdesc) != TYPEFUNC_COMPOSITE)
1967 provsql_error(
"rv_support: expected composite return type");
1968 tupdesc = BlessTupleDesc(tupdesc);
1970 values[0] = Float8GetDatum(iv.first);
1971 values[1] = Float8GetDatum(iv.second);
1973 PG_RETURN_DATUM(HeapTupleGetDatum(heap_form_tuple(tupdesc, values, nulls)));
1974 }
catch (
const std::exception &e) {
ComparisonOperator cmpOpFromOid(Oid op_oid, bool &ok)
Map a PostgreSQL comparison-operator OID to a ComparisonOperator.
AggregationOperator getAggregationOperator(Oid oid)
Map a PostgreSQL aggregate function OID to an AggregationOperator.
Typed aggregation value, operator, and aggregator abstractions.
AggregationOperator
SQL aggregation functions tracked by ProvSQL.
@ COUNT
COUNT(*) or COUNT(expr) → integer.
@ SUM
SUM → integer or float.
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.
static CircuitCache cache
Process-local singleton circuit gate cache.
GenericCircuit getJointCircuit(const std::vector< pg_uuid_t > &tokens, std::vector< gate_t > &gates)
Multi-root variant of getJointCircuit.
Build in-memory circuits from the mmap-backed persistent store.
gate_t
Strongly-typed gate identifier.
Exact conjugate-prior posteriors for observe-evidence circuits.
Per-family polymorphic view over a continuous gate_rv distribution (§F.1 class hierarchy).
Analytical expectation / variance / moment evaluator over RV circuits.
Continuous random-variable helpers (distribution parsing, moments).
Datum rv_support(PG_FUNCTION_ARGS)
SQL: rv_support(token uuid, prov uuid, OUT lo float8, OUT hi float8).
Support-based bound check for continuous-RV comparators.
iterator end()
Past-the-end iterator for the cache.
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 resolveGateToZero(gate_t g)
Replace an arbitrary gate (typically gate_times) by gate_zero.
void resolveCmpToPlusOfKGates(gate_t g, const std::vector< gate_t > &ks)
Replace a gate_cmp by a gate_plus over the given per-row K-gates (the OR of the agg's row-presence in...
bool isCategoricalMixture(gate_t g) const
Test whether g is a categorical-form gate_mixture (the explicit provsql.categorical output).
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...
std::pair< unsigned, unsigned > getInfos(gate_t g) const
Return the integer annotation pair for gate g.
static pg_uuid_t value_gate(const char *prefix, const char *suffix, const char *extra)
The value gate of the text str, at the address named name ("value" followed by str,...
std::pair< double, double > compute_support(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Compute the [lo, hi] support interval of a scalar sub-circuit rooted at root.
static std::optional< ClosedFormShape > truncateShape(const ClosedFormShape &s, double lo, double hi)
Conditional shape after truncating the underlying variable to [lo, hi].
std::optional< ClosedFormShape > matchClosedFormDistribution(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Detect any of the closed-form shapes supported by rv_analytical_curves.
std::variant< TruncatedSingleRv, DiracShape, CategoricalShape, BernoulliMixtureShape > ClosedFormShape
One of the closed-form shapes the analytical-curves payload can render: bare RV (continuous PDF/CDF),...
gate_t lift_conditioning(GenericCircuit &gc, gate_t root, std::optional< gate_t > &event_opt)
Lift conditioning out of a scalar arithmetic expression.
static bool extract_mulinput_value(const GenericCircuit &gc, gate_t mul, double &out)
Same parsing applied to a mulinput's outcome label (categorical).
static bool extract_finite_double(const GenericCircuit &gc, gate_t x, double &out)
Parse a gate_value's extra as a finite float8.
unsigned runReflexiveCmpRewriter(GenericCircuit &gc)
Probability-side pre-pass: rewrite HAVING-style gate_cmp gates that are provably TRUE on the agg's va...
double parseDoubleStrict(const std::string &s)
Strictly parse s as a double.
unsigned runRangeCheck(GenericCircuit &gc)
Run the support-based pruning pass over gc.
bool eventIsProvablyInfeasible(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
True iff the conditioning event is provably infeasible for a bare gate_rv root.
std::unique_ptr< Distribution > makeDistribution(const DistributionSpec &spec)
Construct the per-family Distribution for a parsed spec.
static std::optional< double > shape_mass(const ClosedFormShape &s, double lo, double hi)
Unconditional probability mass of a shape over the interval [lo, hi].
std::optional< DistributionSpec > conjugatePosterior(const GenericCircuit &gc, gate_t target, gate_t evidence)
The exact posterior of target given evidence, as a resolved distribution spec, when the circuit match...
std::optional< DistributionSpec > parse_distribution_spec(const std::string &s)
Parse the on-disk text encoding of a gate_rv distribution.
std::optional< std::pair< double, double > > collectRvConstraints(const GenericCircuit &gc, gate_t event_root, gate_t target_rv)
Walk event_root collecting rv op c constraints on target_rv.
std::optional< DistributionTemplate > parse_distribution_template(const std::string &s)
Parse the on-disk text encoding of a gate_rv distribution, keeping wired (token) parameters as wire r...
std::optional< TruncatedSingleRv > matchTruncatedSingleRv(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Detect a closed-form, optionally-truncated single-RV shape.
double cdfAt(const DistributionSpec &d, double c)
Closed-form CDF for a basic continuous distribution.
unsigned runHavingAlwaysTrueRewriter(GenericCircuit &gc)
Uniform error-reporting macros for ProvSQL.
#define provsql_error(fmt,...)
Report a fatal ProvSQL error and abort the current transaction.
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_case
N-ary guarded selection over scalar (RV) children: wires are [guard_1, value_1, .....
@ 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)
#define PROVSQL_AGG_SCALAR_FLAG
Scalar-aggregation flag, stored in the upper bit of a gate_agg's info2 (whose low 31 bits hold the ag...
C++ utility functions for UUID manipulation.
Bernoulli mixture (gate_mixture with the [p_token, x_token, y_token] shape).
std::shared_ptr< ClosedFormShape > right
std::shared_ptr< ClosedFormShape > left
Categorical distribution over a finite outcome set.
std::vector< std::pair< double, double > > outcomes
(value, mass) pairs
Point mass at a finite scalar value (a gate_value root, or an as_random(c) leaf surfaced as a gate_va...
A closed support interval [lo, hi] (±infinity for unbounded).
Detection result for a closed-form, optionally-truncated single-RV shape.