ProvSQL C/C++ API
Adding support for provenance and uncertainty management to PostgreSQL databases
Loading...
Searching...
No Matches
HybridEvaluator.cpp
Go to the documentation of this file.
1/**
2 * @file HybridEvaluator.cpp
3 * @brief Implementation of the peephole simplifier.
4 * See @c HybridEvaluator.h for the full docstring.
5 */
6#include "HybridEvaluator.h"
7
8#include <cmath>
9#include <limits>
10#include <memory>
11#include <optional>
12#include <stack>
13#include <string>
14#include <unordered_map>
15#include <unordered_set>
16#include <utility>
17#include <vector>
18
19#include "Aggregation.h" // ComparisonOperator, cmpOpFromOid
20#include "AnalyticEvaluator.h" // cdfAt
21#include "distributions/Distribution.h" // makeDistribution, affine, closePlusTerms
22#include "Expectation.h" // evaluateBooleanProbability
23#include "MonteCarloSampler.h" // monteCarloRV, monteCarloScalarSamples
24#include "PivotIntegration.h" // simpsonIntegrate, kSimpsonPanels
25#include "RandomVariable.h" // parse_distribution_spec, double_to_text
26extern "C" {
27#include "provsql_utils.h" // gate_type, provsql_arith_op
28}
29#include <algorithm> // std::sort, std::unique, std::upper_bound
30
31namespace provsql {
32
33namespace {
34
35constexpr double NaN = std::numeric_limits<double>::quiet_NaN();
36
37/* Base @c gate_rv leaves whose sampled value must stay coupled across
38 * distinct comparators. The identity-minting folds (@c try_sum_closure,
39 * @c try_times_scalar_rv, @c try_product_closure, @c try_transform_closure,
40 * @c try_neg_rv) replace their gate with a fresh @c gate_rv and orphan
41 * the base RV they consumed; that mints an independent draw. Sound when
42 * the base RV feeds a single comparator side (its marginal is unchanged),
43 * but WRONG when the same leaf feeds two comparator sides that must see
44 * the same draw -- e.g. a latent parameter b feeding both `0 + b` and
45 * `1 + b`, two perfectly correlated events; folding each into an
46 * independent Normal silently applies the independence approximation.
47 *
48 * The set holds every base RV reachable from two or more sibling subtrees
49 * combined non-additively: the two sides of a @c gate_cmp, or two arms of a
50 * non-additive @c gate_arith combinator (MIN / MAX / PERCENTILE / MINUS /
51 * TIMES / DIV / POW) -- e.g. the shared @c T0 in @c min(T0+T1, T0+T2). A
52 * leaf private to a single such subtree -- or one repeated inside an additive
53 * @c PLUS like `x + x`, which @c try_plus_aggregate folds while preserving the
54 * shared identity -- is absent, so those folds still fire. When a candidate base RV is in this set the fold bails,
55 * leaving the @c gate_arith intact so the island decomposer descends to
56 * the shared leaf and couples the comparators. Set for the duration of
57 * @c runHybridSimplifier; null (empty) elsewhere. */
58const std::unordered_set<gate_t> *g_shared_base_rvs = nullptr;
59
60inline bool is_shared_base_rv(gate_t rv)
61{
62 return g_shared_base_rvs != nullptr && g_shared_base_rvs->count(rv) != 0;
63}
64
65/* Collect the base @c gate_rv leaves reachable from @p start through
66 * @c gate_arith composition and @c gate_rv latent-parameter wires (a
67 * parametric RV couples through its parameter leaves). Mirrors the
68 * cmp-footprint walk but seeded from a single gate, so it can be run
69 * once per comparator wire. */
70void collect_reachable_base_rvs(const GenericCircuit &gc, gate_t start,
71 std::unordered_set<gate_t> &out)
72{
73 std::unordered_set<gate_t> seen;
74 std::stack<gate_t> stk;
75 stk.push(start);
76 while (!stk.empty()) {
77 gate_t g = stk.top(); stk.pop();
78 if (!seen.insert(g).second) continue;
79 auto t = gc.getGateType(g);
80 if (t == gate_rv) {
81 out.insert(g);
82 for (gate_t c : gc.getWires(g)) stk.push(c); /* latent params */
83 continue;
84 }
85 if (t == gate_arith) {
86 for (gate_t c : gc.getWires(g)) stk.push(c);
87 continue;
88 }
89 if (t == gate_mixture && !gc.isCategoricalMixture(g)) {
90 const auto &mw = gc.getWires(g);
91 if (mw.size() == 3) { stk.push(mw[1]); stk.push(mw[2]); }
92 }
93 /* gate_value / other: no base-RV identity below. */
94 }
95}
96
97/**
98 * @brief Try to evaluate a @c gate_arith subtree to a scalar constant.
99 *
100 * Recurses over the @c gate_arith ops, parsing @c gate_value leaves
101 * via @c parseDoubleStrict. Returns @c NaN if any leaf is not a
102 * @c gate_value (or fails to parse), if a binary op has the wrong
103 * arity, or if any arith op is unknown. Successful constants of any
104 * value (including @c 0 and @c NaN-shaped values via division) are
105 * returned as @c double literals; the caller distinguishes
106 * "couldn't fold" from "folded to NaN" via @c std::isnan on the
107 * input gate's children, not on the result. In practice provsql
108 * @c gate_value extras never carry @c NaN, so the @c NaN-as-sentinel
109 * convention is unambiguous.
110 */
111double try_eval_constant(const GenericCircuit &gc, gate_t g)
112{
113 auto t = gc.getGateType(g);
114 if (t == gate_value) {
115 try { return parseDoubleStrict(gc.getExtra(g)); }
116 catch (const CircuitException &) { return NaN; }
117 }
118 if (t != gate_arith) return NaN;
119
120 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
121 const auto &wires = gc.getWires(g);
122 if (wires.empty()) return NaN;
123
124 double first = try_eval_constant(gc, wires[0]);
125 if (std::isnan(first)) return NaN;
126
127 switch (op) {
129 if (wires.size() == 1) return std::round(first);
130 else {
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;
135 }
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;
143 /* The value read as its own type: a double is what these evaluators
144 * already carry, so this one is the value itself. */
145 return wires.size() == 1 ? first : NaN;
147 return wires.size() == 1
148 ? static_cast<double>(static_cast<float>(first)) : NaN;
149 case PROVSQL_ARITH_PLUS: {
150 double r = first;
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;
154 r += v;
155 }
156 return r;
157 }
158 case PROVSQL_ARITH_TIMES: {
159 double r = first;
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;
163 r *= v;
164 }
165 return r;
166 }
167 case PROVSQL_ARITH_MINUS: {
168 if (wires.size() != 2) return NaN;
169 double v = try_eval_constant(gc, wires[1]);
170 if (std::isnan(v)) return NaN;
171 return first - v;
172 }
173 case PROVSQL_ARITH_DIV: {
174 if (wires.size() != 2) return NaN;
175 double v = try_eval_constant(gc, wires[1]);
176 /* A zero divisor has no constant value to fold to: declining it leaves
177 * the DIV gate alone, where folding it would put an infinity in a value
178 * gate for every reader downstream. */
179 if (std::isnan(v) || v == 0.0) return NaN;
180 return first / v;
181 }
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);
187 }
189 if (wires.size() != 1) return NaN;
190 return -first;
191 case PROVSQL_ARITH_MAX: {
192 double r = first;
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;
196 r = std::max(r, v);
197 }
198 return r;
199 }
200 case PROVSQL_ARITH_MIN: {
201 double r = first;
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;
205 r = std::min(r, v);
206 }
207 return r;
208 }
209 case PROVSQL_ARITH_POW: {
210 if (wires.size() != 2) return NaN;
211 double e = try_eval_constant(gc, wires[1]);
212 if (std::isnan(e)) return NaN;
213 /* A domain-violating constant (negative base, non-integer
214 * exponent) folds to NaN, which the NaN-as-sentinel convention
215 * reads as "couldn't fold": the gate stays intact and the
216 * sampler raises its actionable domain error instead of a
217 * silent NaN constant appearing in the circuit. */
218 return std::pow(first, e);
219 }
220 case PROVSQL_ARITH_LN:
221 if (wires.size() != 1) return NaN;
222 /* ln of a negative constant is NaN -> stays unfolded, same as POW. */
223 return std::log(first);
225 if (wires.size() != 1) return NaN;
226 return std::exp(first);
228 /* An order-statistic aggregate over a random member set is never a
229 * constant: leave it for the sampler. */
230 return NaN;
231 }
232 return NaN;
233}
234
235/**
236 * @brief Whether the subtree rooted at @p g contains a @c gate_agg.
237 *
238 * The hybrid simplifier is RV-oriented; aggregate arithmetic
239 * (@c gate_arith over @c gate_agg) is a separate feature whose
240 * comparisons are resolved by the HAVING possible-worlds enumeration,
241 * which must see the original operators to apply the correct (integer
242 * floor vs real) division semantics. Rewrites that are sound for
243 * continuous RVs but not for aggregates (notably the DIV-by-constant to
244 * TIMES-by-reciprocal canonicalisation, which discards integer-division
245 * flooring) consult this to leave aggregate subtrees untouched.
246 */
247bool subtree_contains_agg(const GenericCircuit &gc, gate_t g)
248{
249 std::unordered_set<gate_t> seen;
250 std::stack<gate_t> stk;
251 stk.push(g);
252 while (!stk.empty()) {
253 gate_t cur = stk.top(); stk.pop();
254 if (!seen.insert(cur).second) continue;
255 if (gc.getGateType(cur) == gate_agg) return true;
256 for (gate_t ch : gc.getWires(cur)) stk.push(ch);
257 }
258 return false;
259}
260
261/**
262 * @brief Rewrite @p g in place as a @c gate_value carrying @p c.
263 *
264 * Clears wires and infos; the old children become orphans (no parent
265 * reaches them via @p g anymore). This is the same pattern
266 * @c resolveCmpToBernoulli uses for resolved comparators.
267 */
268void replace_with_value(GenericCircuit &gc, gate_t g, double c)
269{
271}
272
273/**
274 * @brief Test whether wire @p g is a @c gate_value parseable to
275 * scalar @p target (within bit-exact equality).
276 */
277bool is_value_equal_to(const GenericCircuit &gc, gate_t g, double target)
278{
279 if (gc.getGateType(g) != gate_value) return false;
280 try { return parseDoubleStrict(gc.getExtra(g)) == target; }
281 catch (const CircuitException &) { return false; }
282}
283
284/**
285 * @brief Identity-element drop for @c PLUS / @c TIMES.
286 *
287 * - @c PLUS: drop @c gate_value:0 wires. If 0 wires remain, fold to
288 * @c gate_value:0.
289 * - @c TIMES: if any wire is @c gate_value:0, fold to @c gate_value:0
290 * (multiplicative absorber, even if other wires are non-constant).
291 * Otherwise drop @c gate_value:1 wires; if 0 wires remain, fold to
292 * @c gate_value:1.
293 *
294 * Returns @c true if @p g was mutated. After a mutation that leaves
295 * @p g as @c gate_arith, the per-gate fixed-point loop in @c simplify
296 * re-runs the rules: a @c PLUS that had three wires reduced to one
297 * looks the same as the original input to the simplifier, so we just
298 * need to terminate when no rule fires.
299 */
300bool try_identity_drop(GenericCircuit &gc, gate_t g)
301{
302 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
303 auto &wires = gc.getWires(g);
304
305 if (op == PROVSQL_ARITH_PLUS) {
306 std::vector<gate_t> kept;
307 kept.reserve(wires.size());
308 for (gate_t w : wires) {
309 if (!is_value_equal_to(gc, w, 0.0)) kept.push_back(w);
310 }
311 if (kept.size() == wires.size()) return false; /* nothing to drop */
312 if (kept.empty()) {
313 replace_with_value(gc, g, 0.0);
314 return true;
315 }
316 wires = std::move(kept);
317 return true;
318 }
319
320 if (op == PROVSQL_ARITH_TIMES) {
321 for (gate_t w : wires) {
322 if (is_value_equal_to(gc, w, 0.0)) {
323 replace_with_value(gc, g, 0.0);
324 return true;
325 }
326 }
327 std::vector<gate_t> kept;
328 kept.reserve(wires.size());
329 for (gate_t w : wires) {
330 if (!is_value_equal_to(gc, w, 1.0)) kept.push_back(w);
331 }
332 if (kept.size() == wires.size()) return false;
333 if (kept.empty()) {
334 replace_with_value(gc, g, 1.0);
335 return true;
336 }
337 wires = std::move(kept);
338 return true;
339 }
340
341 return false;
342}
343
344/**
345 * @brief Decomposition of a PLUS-wire as @c a*Z + b for the
346 * family sum closure.
347 *
348 * - @c rv_gate == invalid (sentinel @c (gate_t)-1) ⇒ pure constant
349 * wire: contributes @p b to the total mean, 0 to the total
350 * variance, and no RV to the footprint.
351 * - @c rv_gate != invalid ⇒ scalar-multiple-of-normal wire:
352 * contributes @c a*μ + b to the total mean, @c a²σ² to the total
353 * variance, and @p rv_gate to the footprint.
354 */
355struct LinearTerm {
356 gate_t rv_gate; ///< Base gate_rv, or invalid for constants.
357 double a; ///< Scalar multiplier (0 for pure constants).
358 double b; ///< Additive offset (0 for pure RV wires).
359};
360
361constexpr gate_t INVALID_GATE = static_cast<gate_t>(-1);
362
363bool is_invalid(gate_t g) { return g == INVALID_GATE; }
364
365/**
366 * @brief Try to interpret @p g as @c a*Z + b for a single base RV.
367 *
368 * Recognised shapes:
369 * - bare @c gate_rv (any distribution): @c (Z=g, a=1, b=0)
370 * - bare @c gate_value: @c (Z=invalid, a=0, b=value)
371 * - @c arith(NEG, child): negate the child's decomposition
372 * - @c arith(TIMES, value:c, child): scale the child's decomposition
373 * by @c c (and symmetrically @c arith(TIMES, child, value:c)).
374 * Only 2-wire @c TIMES with exactly one @c gate_value side is
375 * recognised; other shapes fall through to "not decomposable".
376 *
377 * Nested @c arith(PLUS, ...) children of the outer PLUS are not
378 * decomposed by this routine: the bottom-up simplifier already
379 * folded them before the outer PLUS is processed, so by the time
380 * we examine the outer PLUS its children are either leaves or
381 * non-foldable arith. An undecomposable wire causes the caller to
382 * bail.
383 *
384 * Distribution-kind concerns are the caller's responsibility:
385 * @c try_sum_closure parses each base RV's spec and dispatches on the
386 * families present via the ClosureRuleRegistry, while
387 * @c try_plus_aggregate is kind-agnostic because the aggregation
388 * rewrite preserves the base-RV identity.
389 */
390std::optional<LinearTerm>
391decompose_linear_term(const GenericCircuit &gc, gate_t g)
392{
393 auto t = gc.getGateType(g);
394
395 if (t == gate_value) {
396 double v;
397 try { v = parseDoubleStrict(gc.getExtra(g)); }
398 catch (const CircuitException &) { return std::nullopt; }
399 return LinearTerm{INVALID_GATE, 0.0, v};
400 }
401
402 if (t == gate_rv) {
403 /* Any RV kind: aggregation only depends on identity, not on
404 * closed-form scaling. The sum closure dispatches on the family
405 * externally. */
406 return LinearTerm{g, 1.0, 0.0};
407 }
408
409 if (t == gate_mixture) {
410 /* A @c gate_mixture (3-wire Bernoulli or categorical N-wire) is a
411 * scalar-RV leaf: two references to the same @c gate_t produce
412 * perfectly-correlated draws of the same RV. Treat it like a
413 * @c gate_rv so the PLUS aggregator can collapse same-mixture
414 * terms (e.g. @c X+X to @c 2·X, @c X-X to @c 0). The in-place
415 * op-change to TIMES then triggers @c try_mixture_lift to push the
416 * scalar inside the branches (3-wire) or the mulinputs'
417 * value text (categorical). The sum closure parses the rv leaf's
418 * spec via @c parse_distribution_spec, which returns @c nullopt on
419 * a mixture's empty extra, so it automatically bails when the
420 * LHS-RV side is a mixture. */
421 return LinearTerm{g, 1.0, 0.0};
422 }
423
424 if (t != gate_arith) return std::nullopt;
425
426 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
427 const auto &wires = gc.getWires(g);
428
429 /* After an identity-element drop, a PLUS or TIMES gate can be left
430 * with a single wire that semantically passes through. Recurse so
431 * the outer closure can still see the underlying term. We can't
432 * fold the singleton wrapper away in place (rewriting it as the
433 * child's type / extra would mint a fresh RV identity and break
434 * per-iteration MC memoisation across other parents of the child),
435 * but the outer closure rewrites the OUTER gate, which is safe. */
436 if ((op == PROVSQL_ARITH_PLUS || op == PROVSQL_ARITH_TIMES)
437 && wires.size() == 1) {
438 return decompose_linear_term(gc, wires[0]);
439 }
440
441 if (op == PROVSQL_ARITH_NEG) {
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};
446 }
447
448 if (op == PROVSQL_ARITH_TIMES) {
449 if (wires.size() != 2) return std::nullopt;
450 /* Identify the constant side and the variable side. */
451 double c = NaN;
452 gate_t var_side = INVALID_GATE;
453 if (gc.getGateType(wires[0]) == gate_value) {
454 try { c = parseDoubleStrict(gc.getExtra(wires[0])); }
455 catch (const CircuitException &) { return std::nullopt; }
456 var_side = wires[1];
457 } else if (gc.getGateType(wires[1]) == gate_value) {
458 try { c = parseDoubleStrict(gc.getExtra(wires[1])); }
459 catch (const CircuitException &) { return std::nullopt; }
460 var_side = wires[0];
461 } else {
462 return std::nullopt;
463 }
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};
467 }
468
469 return std::nullopt;
470}
471
472/**
473 * @brief Family closure on a @c PLUS gate, driven by the
474 * @c ClosureRuleRegistry.
475 *
476 * Decomposes every wire to @c a*Z + b (via @c decompose_linear_term),
477 * parses each base RV's distribution, and hands the terms to
478 * @c closePlusTerms, which dispatches on the families present. The
479 * registered rules cover:
480 *
481 * - Normal: any linear combination of independent normals (plus
482 * constants) folds to a single normal;
483 * - Exponential / Erlang: an unscaled same-rate chain folds to
484 * Erlang(Σk, λ) -- left-associative parsing of <tt>a + b + c</tt>
485 * builds <tt>(a+b)+c</tt> which bottom-up simplifies to
486 * Erlang(2)+c, so the rule accepts the mixed Erlang+Exp shape to
487 * close the chain;
488 * - Uniform: a single (possibly scaled / negated) uniform plus
489 * constants folds to the affine-transformed uniform, including the
490 * post-MINUS-canonicalisation shapes @c c + (-U) and @c (-U) + c.
491 * @c U + @c U is @b not closed (triangular density), which the rule
492 * expresses by declining a second Uniform term.
493 *
494 * Independence is tested here, structurally: every non-constant term
495 * must have a distinct base-RV @c gate_t (each RV constructor mints a
496 * fresh UUID, so distinctness implies independence, and
497 * @c try_plus_aggregate runs first so shared-UUID terms were already
498 * consolidated). A @c gate_mixture leaf has no parseable distribution
499 * spec, so mixture-bearing sums bail (they are @c try_mixture_lift's
500 * job). When every wire is a pure constant the dispatch declines and
501 * the constant fold handles the gate on the next fixed-point iteration.
502 *
503 * Same coupling caveat as @c try_times_scalar_rv: replacing @p g with
504 * a fresh @c gate_rv mints a new RV identity.
505 */
506bool try_sum_closure(GenericCircuit &gc, gate_t g)
507{
508 const auto &wires = gc.getWires(g);
509 if (wires.size() < 2) return false;
510
511 std::vector<LinearTerm> lterms;
512 lterms.reserve(wires.size());
513 for (gate_t w : wires) {
514 auto term = decompose_linear_term(gc, w);
515 if (!term) return false;
516 lterms.push_back(*term);
517 }
518
519 /* Independence test + per-term distribution parse. */
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});
528 continue;
529 }
530 if (!seen_rvs.insert(t.rv_gate).second) return false; /* dependent */
531 /* A base RV shared with sibling subtrees must stay a live wire so
532 * downstream coupling survives; folding it into a fresh identity
533 * here would decouple the correlated events. */
534 if (is_shared_base_rv(t.rv_gate)) return false;
535 auto spec = parse_distribution_spec(gc.getExtra(t.rv_gate));
536 if (!spec) return false; /* mixture / corrupted extra */
537 dists[i] = makeDistribution(*spec);
538 terms.push_back({dists[i].get(), t.a, t.b});
539 }
540
541 auto folded = closePlusTerms(terms);
542 if (!folded) return false;
543
544 gc.resolveToRv(g, folded->serialise());
545 return true;
546}
547
548/**
549 * @brief Product closure on a @c TIMES gate, driven by the
550 * @c ProductRuleRegistry.
551 *
552 * Wires must be @c gate_value factors (multiplied into one scalar) or
553 * bare @c gate_rv leaves with distinct UUIDs (independence, as in
554 * @c try_sum_closure); at least two RV factors, or the 2-wire
555 * scalar-times-RV shape is @c try_times_scalar_rv's job. The
556 * registered rules cover lognormal products (parameters add in log
557 * space); the accumulated scalar then applies through the family's
558 * @c affine. Same fresh-identity coupling caveat as the sum closure.
559 */
560bool try_product_closure(GenericCircuit &gc, gate_t g)
561{
562 const auto &wires = gc.getWires(g);
563 if (wires.size() < 2) return false;
564
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;
569 for (gate_t w : wires) {
570 const auto t = gc.getGateType(w);
571 if (t == gate_value) {
572 try { c_total *= parseDoubleStrict(gc.getExtra(w)); }
573 catch (const CircuitException &) { return false; }
574 continue;
575 }
576 if (t != gate_rv) return false;
577 if (!seen_rvs.insert(w).second) return false; /* dependent */
578 if (is_shared_base_rv(w)) return false; /* shared: keep live */
579 auto spec = parse_distribution_spec(gc.getExtra(w));
580 if (!spec) return false;
581 dists.push_back(makeDistribution(*spec));
582 factors.push_back(dists.back().get());
583 }
584 if (factors.size() < 2) return false;
585
586 auto combined = closeProductFactors(factors);
587 if (!combined) return false;
588 if (c_total != 1.0) {
589 combined = combined->scale(c_total);
590 if (!combined) return false;
591 }
592 gc.resolveToRv(g, combined->serialise());
593 return true;
594}
595
596/**
597 * @brief Transform closure on a unary @c LN / @c EXP gate, driven by
598 * the @c TransformRuleRegistry.
599 *
600 * When the child is a bare @c gate_rv whose family registers a
601 * closed-form image (exp(normal) is lognormal, ln(lognormal) is
602 * normal), the gate folds to the image distribution -- the bottom-up
603 * pass has already folded the child, so chains like
604 * <tt>exp(normal + normal)</tt> collapse fully. Same fresh-identity
605 * coupling caveat as the sum closure.
606 */
607bool try_transform_closure(GenericCircuit &gc, gate_t g)
608{
609 const auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
610 const char *transform = op == PROVSQL_ARITH_LN ? "ln"
611 : op == PROVSQL_ARITH_EXP ? "exp"
612 : nullptr;
613 if (!transform) return false;
614 const auto &wires = gc.getWires(g);
615 if (wires.size() != 1) return false;
616 if (gc.getGateType(wires[0]) != gate_rv) return false;
617 if (is_shared_base_rv(wires[0])) return false; /* shared: keep live */
618 auto spec = parse_distribution_spec(gc.getExtra(wires[0]));
619 if (!spec) return false;
620
621 auto image = closeTransform(transform, *makeDistribution(*spec));
622 if (!image) return false;
623 gc.resolveToRv(g, image->serialise());
624 return true;
625}
626
627/**
628 * @brief Negation closure on a bare @c gate_rv: rewrite @c arith(NEG, Z)
629 * as a closed-form-negated @c gate_rv when @c Z's family admits
630 * one.
631 *
632 * Delegates to @c Distribution::negate (@c affine(-1, 0)): Normal and
633 * Uniform fold (<tt>-N(μ, σ) = N(-μ, σ)</tt>,
634 * <tt>-U(a, b) = U(-b, -a)</tt>); Exponential / Erlang decline (the
635 * support flips to @c (-∞, 0], leaving the family).
636 *
637 * Coupling discipline: same as @c try_times_scalar_rv. Pass-2 gated
638 * so a parent PLUS containing @c NEG(Z) and a sibling reference to the
639 * same @c Z is folded first by @c try_plus_aggregate (which recognises
640 * @c NEG via @c decompose_linear_term's coefficient @c -1) before we
641 * mint a fresh @c gate_rv at the NEG.
642 */
643bool try_neg_rv(GenericCircuit &gc, gate_t g)
644{
645 if (gc.getGateType(g) != gate_arith) return false;
646 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
647 if (op != PROVSQL_ARITH_NEG) return false;
648 const auto &wires = gc.getWires(g);
649 if (wires.size() != 1) return false;
650 if (gc.getGateType(wires[0]) != gate_rv) return false;
651 if (is_shared_base_rv(wires[0])) return false; /* shared: keep live */
652
653 auto spec = parse_distribution_spec(gc.getExtra(wires[0]));
654 if (!spec) return false;
655
656 auto negated = makeDistribution(*spec)->negate();
657 if (!negated) return false;
658 gc.resolveToRv(g, negated->serialise());
659 return true;
660}
661
662/**
663 * @brief Mixture-lift rewrite: push @c PLUS / @c TIMES inside a
664 * single @c gate_mixture child.
665 *
666 * Fires on a @c gate_arith with op @c PLUS or @c TIMES whose children
667 * contain exactly one @c gate_mixture. Replaces the parent with a
668 * @c gate_mixture sharing the same Bernoulli (so the original
669 * <tt>p_token</tt> identity is preserved and any other gate that
670 * referenced it continues to see it):
671 *
672 * <tt>a + mixture(p, X, Y) → mixture(p, a + X, a + Y)</tt>
673 *
674 * The two new branches are fresh @c gate_arith children built via
675 * @c addAnonymousArithGate; each is then re-fed to @c apply_rules so
676 * the family sum closure gets a chance
677 * to collapse them. This is the source of the headline simplifier
678 * gain for compound RV expressions: <tt>3 + mixture(p, N(0,1), N(2,1))</tt>
679 * folds to <tt>mixture(p, N(3,1), N(5,1))</tt> in a single bottom-up
680 * pass.
681 *
682 * Multi-mixture lifts (two or more @c gate_mixture children of the
683 * same arith) are out of scope: each would multiply the branch count
684 * by 2 and the lifted form would couple the resulting branches
685 * through their Bernoullis, which the current closures cannot
686 * collapse further. @c MINUS / @c DIV / @c NEG lifts are also out of
687 * scope (the user requested only @c PLUS and @c TIMES); they can be
688 * added in a follow-up once the sum closure handles
689 * subtraction.
690 *
691 * Returns @c true if @p g was mutated.
692 */
693unsigned apply_rules(GenericCircuit &gc, gate_t g,
694 bool include_scalar_fold); /* forward decl */
695
696/**
697 * @brief Categorical-mixture lift helper.
698 *
699 * Pushes a constant scaling (@c TIMES) or offset (@c PLUS) inside the
700 * N-wire categorical-form @c gate_mixture <tt>[key, mul_1, ..., mul_n]</tt>
701 * by minting a fresh categorical mixture sharing the same @p key gate
702 * and one new @c gate_mulinput per outcome with an updated value text.
703 *
704 * Sharing the key preserves the semantic that the new mixture is a
705 * deterministic function of the same underlying categorical draw (so
706 * <tt>c · X</tt> and @c X stay perfectly correlated downstream via
707 * FootprintCache key-overlap dependency tracking). All other arith
708 * wires must be @c gate_value constants; an RV factor / offset cannot
709 * be pushed into a mulinput's scalar @c extra so the rule bails.
710 *
711 * Returns @c true if @p g was mutated.
712 */
713bool try_categorical_mixture_lift(GenericCircuit &gc, gate_t g,
715 gate_t mix_gate,
716 const std::vector<gate_t> &others)
717{
718 if (op != PROVSQL_ARITH_PLUS && op != PROVSQL_ARITH_TIMES) return false;
719
720 /* Combine the non-mixture wires into a single scalar offset (PLUS)
721 * or factor (TIMES). Bail on any non-value wire: an RV factor /
722 * offset cannot be pushed into a mulinput's value text. */
723 double offset = 0.0;
724 double factor = 1.0;
725 for (gate_t w : others) {
726 if (gc.getGateType(w) != gate_value) return false;
727 double v;
728 try { v = parseDoubleStrict(gc.getExtra(w)); }
729 catch (const CircuitException &) { return false; }
730 if (op == PROVSQL_ARITH_PLUS) offset += v;
731 else factor *= v;
732 }
733
734 /* Build the new wire list: same key (preserves correlation with the
735 * original categorical) and one fresh mulinput per outcome with the
736 * transformed value text. Snapshot the mixture's wires by value:
737 * @c addAnonymousMulinputGateWithValue below calls @c addGate, which
738 * does @c wires.push_back({}) on the circuit's outer wire vector,
739 * and that can reallocate -- invalidating any reference returned by
740 * @c getWires. Reads of the reference after the first iteration
741 * then return garbage gate ids, which surfaces either as wrong
742 * outcome values or as a backend crash. */
743 const std::vector<gate_t> mw = gc.getWires(mix_gate);
744 const gate_t key = mw[0];
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];
750 double old_v;
751 try { old_v = parseDoubleStrict(gc.getExtra(old_mul)); }
752 catch (const CircuitException &) { return false; }
753 const double new_v = (op == PROVSQL_ARITH_PLUS)
754 ? (offset + old_v)
755 : (factor * old_v);
756 const double p = gc.getProb(old_mul);
757 const auto vi = static_cast<unsigned>(gc.getInfos(old_mul).first);
759 key, p, vi, double_to_text(new_v));
760 new_wires.push_back(new_mul);
761 }
762 gc.resolveToCategoricalMixture(g, std::move(new_wires));
763 return true;
764}
765
766bool try_mixture_lift(GenericCircuit &gc, gate_t g,
767 bool include_scalar_fold)
768{
769 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
770 if (op != PROVSQL_ARITH_PLUS && op != PROVSQL_ARITH_TIMES) return false;
771
772 const auto &wires = gc.getWires(g);
773 if (wires.size() < 2) return false; /* nothing to lift */
774
775 /* Find exactly one mixture child. */
776 std::size_t mix_idx = static_cast<std::size_t>(-1);
777 for (std::size_t i = 0; i < wires.size(); ++i) {
778 if (gc.getGateType(wires[i]) == gate_mixture) {
779 if (mix_idx != static_cast<std::size_t>(-1)) return false;
780 mix_idx = i;
781 }
782 }
783 if (mix_idx == static_cast<std::size_t>(-1)) return false;
784
785 const auto mix_gate = wires[mix_idx];
786
787 /* Snapshot the remaining wires. We need a copy because the
788 * resolveToMixture / resolveToCategoricalMixture calls below clear
789 * the parent's wire vector. */
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]);
794 }
795
796 /* Categorical N-wire form: push the constant offset / factor into
797 * each mulinput's value text. RV factors / offsets cannot be pushed
798 * into mulinput leaves so the rule bails on those. */
799 if (gc.isCategoricalMixture(mix_gate)) {
800 return try_categorical_mixture_lift(gc, g, op, mix_gate, others);
801 }
802
803 /* Classic 3-wire Bernoulli mixture. */
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];
809
810 /* Build two new arith children: one with x in the mixture slot,
811 * one with y. Order matters for non-commutative ops, but PLUS /
812 * TIMES are both commutative so we just append the branch RV to
813 * the others. */
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);
816 gate_t new_x = gc.addAnonymousArithGate(op, std::move(new_x_wires));
817 gate_t new_y = gc.addAnonymousArithGate(op, std::move(new_y_wires));
818
819 /* Rewrite g as gate_mixture(p, new_x, new_y). This clears g's
820 * old wires / infos / extra and installs the new structure. */
821 gc.resolveToMixture(g, p_tok, new_x, new_y);
822
823 /* Recursively fold the two new arith children so they get a chance
824 * to collapse via the family sum closure. Each is
825 * itself a gate_arith of the same op, with at least 2 wires (the
826 * "others" we copied plus the branch RV), so apply_rules's
827 * PLUS/TIMES path is the correct entry point. The scalar-fold flag
828 * is propagated so pass-2's scalar-times-RV closure stays the only
829 * place that mints a fresh @c gate_rv at a scaled-RV TIMES site
830 * (avoids losing shared-RV identity in front of a sibling PLUS). */
831 apply_rules(gc, new_x, include_scalar_fold);
832 apply_rules(gc, new_y, include_scalar_fold);
833
834 return true;
835}
836
837/**
838 * @brief Scalar-times-RV closure: fold @c arith(TIMES, value:c, Z) to
839 * a single closed-form-scaled @c gate_rv.
840 *
841 * Fires on a 2-wire @c TIMES whose wires are exactly one @c gate_value
842 * (the scalar @c c) and one @c gate_rv leaf @c Z whose distribution
843 * admits a closed-form scale transform, per @c Distribution::scale
844 * (@c affine(c, 0)): Normal for any non-zero @c c, Uniform for any
845 * non-zero @c c (a negative @c c flips the bounds), Exponential /
846 * Erlang for @c c > 0 only (negative scaling flips the support).
847 *
848 * The c=0 absorber and c=1 identity are handled by
849 * @c try_identity_drop, so this rule defensively bails on them to
850 * avoid a duplicate rewrite path. RV kinds without a closed-form
851 * scaling fall through.
852 *
853 * Coupling caveat (shared with @c try_sum_closure): replacing the
854 * TIMES with a fresh @c gate_rv mints a new RV identity at @p g, so
855 * any other path that references @c Z and shares a downstream consumer
856 * with @p g will see decoupled draws after the fold. In practice the
857 * rewrite path produces per-row orphan subtrees, so this is consistent
858 * with the family sum closure's behaviour.
859 *
860 * Returns @c true if @p g was mutated.
861 */
862bool try_times_scalar_rv(GenericCircuit &gc, gate_t g)
863{
864 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
865 if (op != PROVSQL_ARITH_TIMES) return false;
866 const auto &wires = gc.getWires(g);
867 if (wires.size() != 2) return false;
868
869 /* Identify the value side and the rv side. */
870 double c = NaN;
871 gate_t rv_side = INVALID_GATE;
872 if (gc.getGateType(wires[0]) == gate_value
873 && gc.getGateType(wires[1]) == gate_rv) {
874 try { c = parseDoubleStrict(gc.getExtra(wires[0])); }
875 catch (const CircuitException &) { return false; }
876 rv_side = wires[1];
877 } else if (gc.getGateType(wires[1]) == gate_value
878 && gc.getGateType(wires[0]) == gate_rv) {
879 try { c = parseDoubleStrict(gc.getExtra(wires[1])); }
880 catch (const CircuitException &) { return false; }
881 rv_side = wires[0];
882 } else {
883 return false;
884 }
885
886 /* c=0 / c=1 are the identity-drop's job; bailing here keeps the
887 * two rules' responsibilities disjoint. */
888 if (c == 0.0 || c == 1.0) return false;
889
890 if (is_shared_base_rv(rv_side)) return false; /* shared: keep live */
891
892 auto spec = parse_distribution_spec(gc.getExtra(rv_side));
893 if (!spec) return false;
894
895 auto scaled = makeDistribution(*spec)->scale(c);
896 if (!scaled) return false;
897
898 /* Defensive: a zero-σ normal collapses to a Dirac. σ=0 normals
899 * are normally constructed via @c as_random by @c provsql.normal,
900 * but if one slipped through (e.g. a future closure produced
901 * σ=0 from the linear combination), route it through value. */
902 if (auto dirac = scaled->asDirac()) {
903 replace_with_value(gc, g, *dirac);
904 return true;
905 }
906
907 gc.resolveToRv(g, scaled->serialise());
908 return true;
909}
910
911/**
912 * @brief PLUS coefficient aggregation: collapse same-base-RV terms
913 * in a sum.
914 *
915 * For a @c PLUS gate whose every wire decomposes via
916 * @c decompose_linear_term to <tt>a·Z + b</tt>, sums the coefficients
917 * per @c rv_gate UUID and accumulates all the constant offsets into a
918 * single @c b_total. Rebuilds the wire list as one @c TIMES per
919 * surviving RV (or a bare RV wire when its coefficient is exactly @c 1)
920 * plus a single @c value wire for @c b_total when non-zero.
921 *
922 * Fires when at least one of the following holds:
923 * - some @c rv_gate appears in more than one wire (the X+X case);
924 * - more than one constant wire is present (consolidates them).
925 *
926 * Without these triggers the rebuild would be a no-op or worse
927 * (minting fresh @c TIMES wrappers identical in shape to existing
928 * input wires), so the rule bails to keep the simplifier idempotent.
929 *
930 * Unlike @c try_sum_closure / @c try_times_scalar_rv, this rule is
931 * @b safe under shared base-RV identity: the rebuild preserves every
932 * @c rv_gate as a wire (wrapped in @c TIMES when its coefficient is
933 * non-unit), so any other path that referenced @c Z continues to see
934 * the same gate. The subsequent fold of <tt>arith(TIMES, value:a, Z)</tt>
935 * by @c try_times_scalar_rv inherits the same coupling caveat as the
936 * family sum closure (see its docstring).
937 *
938 * Returns @c true if @p g was mutated.
939 */
940bool try_plus_aggregate(GenericCircuit &gc, gate_t g,
941 bool include_scalar_fold)
942{
943 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
944 if (op != PROVSQL_ARITH_PLUS) return false;
945 const auto &wires_in = gc.getWires(g);
946 if (wires_in.size() < 2) return false;
947
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;
953 terms.push_back(*t);
954 }
955
956 /* Aggregate per rv_gate. A vector preserves insertion order so the
957 * rebuilt wire list is deterministic across runs; the per-PLUS
958 * arity is small enough that O(n²) lookup is fine. */
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) {
963 b_total += t.b;
964 if (is_invalid(t.rv_gate)) {
965 ++constants_in;
966 continue;
967 }
968 bool found = false;
969 for (auto &p : coeffs) {
970 if (p.first == t.rv_gate) {
971 p.second += t.a;
972 found = true;
973 break;
974 }
975 }
976 if (!found) coeffs.emplace_back(t.rv_gate, t.a);
977 }
978
979 /* Fire only when there's actual consolidation to do. Without a
980 * duplicate RV (or multiple constants) the rebuild would mint
981 * shape-equivalent TIMES wrappers for input wires like
982 * arith(TIMES, value:a, Z), oscillating the gate vector. */
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;
986
987 /* Drop zero-coefficient RVs (X + (-X) survivors). */
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);
992 }
993
994 /* All RVs canceled: fold g to a value gate carrying b_total. */
995 if (kept.empty()) {
996 replace_with_value(gc, g, b_total);
997 return true;
998 }
999
1000 /* Single surviving RV term with no constant offset. Rewrite g
1001 * directly in place as the simplest representation:
1002 * - a == 1 ⇒ singleton PLUS([Z]) (we can't safely dissolve to Z
1003 * in place because that would mint a fresh RV identity at g).
1004 * - a != 1 ⇒ in-place op-change from PLUS to TIMES with wires
1005 * [value:a, Z]. When @p include_scalar_fold is set the fixed-point
1006 * loop then re-enters apply_rules on g (now a TIMES), giving
1007 * try_times_scalar_rv a chance to fold the scaled RV. Pass 1
1008 * runs with @p include_scalar_fold = false (deferring the fold so
1009 * the outer aggregator sees @c c·X-shaped children with intact
1010 * RV identity); pass 2 then folds the surviving TIMES wrapper.
1011 * Either way, the in-place op-change avoids the PLUS([TIMES(..)])
1012 * double wrapper that would otherwise hide the bare-RV shape from
1013 * @c AnalyticEvaluator's @c bareRv lookup. */
1014 if (kept.size() == 1 && b_total == 0.0) {
1015 const auto &only = kept.front();
1016 if (only.second == 1.0) {
1017 gc.setWires(g, {only.first});
1018 } else {
1019 const gate_t cv = gc.addAnonymousValueGate(
1020 double_to_text(only.second));
1021 gc.setInfos(g, static_cast<unsigned>(PROVSQL_ARITH_TIMES), 0);
1022 gc.setWires(g, {cv, only.first});
1023 }
1024 return true;
1025 }
1026
1027 /* General case: rebuild g as a multi-wire PLUS. */
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);
1033 } else {
1034 const gate_t cv = gc.addAnonymousValueGate(double_to_text(p.second));
1036 {cv, p.first});
1037 new_wires.push_back(tm);
1038 }
1039 }
1040 if (b_total != 0.0) {
1041 new_wires.push_back(gc.addAnonymousValueGate(double_to_text(b_total)));
1042 }
1043
1044 gc.setWires(g, std::move(new_wires));
1045
1046 /* Recurse into freshly-minted TIMES children so try_times_scalar_rv
1047 * gets a chance to fold them within the same bottom-up pass when
1048 * @p include_scalar_fold is set. Same pattern as try_mixture_lift. */
1049 for (gate_t w : gc.getWires(g)) {
1050 if (gc.getGateType(w) == gate_arith) {
1051 apply_rules(gc, w, include_scalar_fold);
1052 }
1053 }
1054 return true;
1055}
1056
1057/**
1058 * @brief Run the per-gate fixed-point loop.
1059 *
1060 * After each rule succeeds the gate is re-evaluated under every rule,
1061 * so a single bottom-up pass collapses nested foldable structures
1062 * (e.g. <tt>arith(NEG, arith(PLUS, value, value))</tt>) in one go.
1063 *
1064 * @return Number of rewrites performed on this gate.
1065 */
1066unsigned apply_rules(GenericCircuit &gc, gate_t g,
1067 bool include_scalar_fold)
1068{
1069 unsigned local = 0;
1070 /* Iteration bound: each rule strictly shrinks the gate (fewer wires
1071 * or simpler type), so the loop terminates in O(#initial wires)
1072 * iterations. The bound is defensive insurance against an
1073 * unintended infinite loop. */
1074 for (unsigned iter = 0; iter < 32; ++iter) {
1075 if (gc.getGateType(g) != gate_arith) break;
1076
1077 /* 1. Constant folding (collapses any all-gate_value arith). */
1078 {
1079 double c = try_eval_constant(gc, g);
1080 if (!std::isnan(c)) {
1081 replace_with_value(gc, g, c);
1082 ++local;
1083 break;
1084 }
1085 }
1086
1087 /* 1b. MINUS-to-PLUS canonicalisation. Rewrites
1088 * @c arith(MINUS, A, B) as @c arith(PLUS, A, arith(NEG, B))
1089 * so every downstream rule -- PLUS aggregation, family
1090 * closures, mixture-lift -- only needs to handle PLUS.
1091 * @c decompose_linear_term already recognises @c NEG as a
1092 * coefficient @c -1, so the rewritten parent's
1093 * @c decompose_linear_term yields the same linear-term shape
1094 * as the original MINUS would have, modulo one extra
1095 * gate_arith level for the NEG. Runs after constant fold so
1096 * a fully-constant @c MINUS(value, value) collapses to a
1097 * @c value gate without minting an interim NEG. */
1098 {
1099 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
1100 if (op == PROVSQL_ARITH_MINUS) {
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];
1106 {b});
1107 gc.setInfos(g, static_cast<unsigned>(PROVSQL_ARITH_PLUS), 0);
1108 gc.setWires(g, {a, neg_b});
1109 ++local;
1110 continue;
1111 }
1112 }
1113 }
1114
1115 /* 1c. DIV-by-constant to TIMES-by-reciprocal canonicalisation.
1116 * Rewrites @c arith(DIV, X, value:c) as
1117 * @c arith(TIMES, X, value:1/c) (c != 0) so the existing
1118 * scalar-times-RV closure (@c try_times_scalar_rv) and every
1119 * other downstream TIMES rule fold @c X/c uniformly with
1120 * @c c*X. DIV-by-non-constant is left alone (no closure to
1121 * apply); fully-constant @c DIV(value, value) is handled by
1122 * the constant fold above so we never see @c c=0 here.
1123 * Aggregate divisions (an @c X bearing a @c gate_agg) are left
1124 * intact: their HAVING possible-worlds enumeration applies the
1125 * correct integer-floor / real division on the original DIV,
1126 * which a TIMES-by-reciprocal would silently discard. */
1127 {
1128 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
1129 if (op == PROVSQL_ARITH_DIV) {
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];
1135 const gate_t inv = gc.addAnonymousValueGate(
1136 double_to_text(1.0 / c));
1137 gc.setInfos(g, static_cast<unsigned>(PROVSQL_ARITH_TIMES), 0);
1138 gc.setWires(g, {x, inv});
1139 ++local;
1140 continue;
1141 }
1142 }
1143 }
1144 }
1145
1146 /* 2. Identity / absorber drops on PLUS and TIMES. */
1147 if (try_identity_drop(gc, g)) {
1148 ++local;
1149 continue;
1150 }
1151
1152 /* 3. Mixture lift: push PLUS / TIMES inside a single mixture
1153 * child. Runs BEFORE the normal / erlang closures so the
1154 * branch arith children get to try those closures themselves
1155 * after the lift. Once the lift fires the parent is no
1156 * longer gate_arith, so the loop terminates on the next
1157 * iteration via the gate_arith guard above. */
1158 if (try_mixture_lift(gc, g, include_scalar_fold)) {
1159 ++local;
1160 break;
1161 }
1162
1163 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
1164
1165 /* 4. PLUS coefficient aggregation: collapse X+X, X-X, multiple
1166 * constants, etc. Runs BEFORE the family closures so they see
1167 * a sum with distinct RV identities (which they assume), and
1168 * so X+X folds through the scalar-times-RV closure on the
1169 * minted 2*X child. */
1170 if (op == PROVSQL_ARITH_PLUS) {
1171 if (try_plus_aggregate(gc, g, include_scalar_fold)) {
1172 ++local;
1173 continue;
1174 }
1175 }
1176
1177 /* 5. Scalar-times-RV closure on TIMES: c · gate_rv folds to a
1178 * closed-form-scaled gate_rv for the supported families. Gated
1179 * by @p include_scalar_fold: the bottom-up DFS visits children
1180 * before parents, and folding @c c·X to a fresh @c gate_rv at
1181 * the TIMES gate would lose @c X's identity, which an outer
1182 * @c PLUS-aggregation sibling like @c x in @c 2·x+x relies on
1183 * to recognise the shared base RV. Pass 1 runs all other rules
1184 * so the aggregator gets first crack at @c c·X-shaped wires;
1185 * pass 2 then folds the remaining TIMES gates with this rule
1186 * via @c runHybridSimplifier's post-pass. */
1187 if (op == PROVSQL_ARITH_TIMES && include_scalar_fold) {
1188 if (try_times_scalar_rv(gc, g)) {
1189 ++local;
1190 break;
1191 }
1192 }
1193
1194 /* 6. Family closures, dispatched on the families present through
1195 * the closure registries:
1196 * - PLUS: normal linear combinations, same-rate Exp/Erlang
1197 * chains, single-Uniform affine shapes;
1198 * - TIMES: lognormal products (parameters add in log space);
1199 * - LN / EXP: the normal <-> lognormal transform bridges. */
1200 if (op == PROVSQL_ARITH_PLUS) {
1201 if (try_sum_closure(gc, g)) { ++local; break; }
1202 }
1203 if (op == PROVSQL_ARITH_TIMES) {
1204 if (try_product_closure(gc, g)) { ++local; break; }
1205 }
1206 if (op == PROVSQL_ARITH_LN || op == PROVSQL_ARITH_EXP) {
1207 if (try_transform_closure(gc, g)) { ++local; break; }
1208 }
1209
1210 break; /* no rule fired this iteration */
1211 }
1212 return local;
1213}
1214
1215/**
1216 * @brief Post-order DFS that simplifies every reachable gate.
1217 *
1218 * Children are simplified before parents so by the time a gate is
1219 * examined its wires already reflect any rewrites: the bottom-up
1220 * order is essential for cascading folds (a parent PLUS over a child
1221 * arith that just folded to a gate_value gets a chance to fold that
1222 * constant away).
1223 */
1224void simplify(GenericCircuit &gc, gate_t g,
1225 std::unordered_set<gate_t> &done, unsigned &counter,
1226 bool include_scalar_fold)
1227{
1228 /* Iterative DFS with an explicit stack: the natural recursive form
1229 * blew the host stack on deeply-nested arith chains in early
1230 * experiments; iteration with a small per-node bookkeeping triple
1231 * (gate, child-cursor, processed-flag) keeps the cost in heap. */
1232 std::stack<std::pair<gate_t, std::size_t>> stk;
1233 if (!done.insert(g).second) return;
1234 stk.emplace(g, 0);
1235
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);
1243 continue;
1244 }
1245 /* All children processed; apply rules to cur. */
1246 if (gc.getGateType(cur) == gate_arith)
1247 counter += apply_rules(gc, cur, include_scalar_fold);
1248 stk.pop();
1249 }
1250}
1251
1252} // namespace
1253
1255{
1256 unsigned counter = 0;
1257 /* Walk every gate in order: @c try_eval_constant recurses through
1258 * @c gate_arith children itself (via @c try_eval_constant's own
1259 * recursion on @c gate_arith ops + base case at @c gate_value),
1260 * so a single linear pass over the gate indices is sufficient.
1261 * No DFS bookkeeping needed because the rewrite produces a
1262 * @c gate_value (terminal), never another @c gate_arith. */
1263 const auto nb = gc.getNbGates();
1264 for (std::size_t i = 0; i < nb; ++i) {
1265 auto g = static_cast<gate_t>(i);
1266 if (gc.getGateType(g) != gate_arith) continue;
1267 double c = try_eval_constant(gc, g);
1268 if (!std::isnan(c)) {
1269 replace_with_value(gc, g, c);
1270 ++counter;
1271 }
1272 }
1273 return counter;
1274}
1275
1277{
1278 unsigned counter = 0;
1279 const auto nb = gc.getNbGates();
1280 for (std::size_t i = 0; i < nb; ++i) {
1281 auto g = static_cast<gate_t>(i);
1282 if (gc.getGateType(g) != gate_mixture) continue;
1283 /* Categorical N-wire mixtures carry their masses in the mulinput
1284 * wires, not a single Bernoulli selector; the degenerate collapse
1285 * below is the classic 3-wire Bernoulli shape only. */
1286 if (gc.isCategoricalMixture(g)) continue;
1287 const auto &wires = gc.getWires(g);
1288 if (wires.size() != 3) continue;
1289
1290 /* The mixing weight pi = P(selector = true). We fold only when pi
1291 * is known to be exactly 1 or 0 without a probability computation:
1292 * a resolved identity selector (gate_one / gate_zero) or a bare
1293 * Bernoulli gate_input whose pinned probability is 1 or 0. In the
1294 * loaded circuit an un-set_prob'd input carries the default
1295 * probability 1 (GenericCircuit::setGate), which is exactly the
1296 * "provenance-tracked but not probabilistic" tuple: certainly
1297 * present. A compound Boolean selector is left intact -- its pi
1298 * would require a (possibly #P-hard) probability evaluation and
1299 * depends on mutable input probabilities, so it belongs to the
1300 * probability-aware evaluators, not this universal load-time pass. */
1301 const gate_t sel = wires[0];
1302 double pi;
1303 switch (gc.getGateType(sel)) {
1304 case gate_one: pi = 1.0; break;
1305 case gate_zero: pi = 0.0; break;
1306 case gate_input: pi = gc.getProb(sel); break;
1307 default: continue;
1308 }
1309
1310 /* A degenerate Bernoulli carries no coupling: collapsing the
1311 * mixture to its surviving arm is exact even when the selector is
1312 * shared with other mixtures (they too are deterministic). This is
1313 * unsound at any fractional pi, which is why the switch above admits
1314 * only the exact 0/1 cases. liftConditionedToTarget rewrites g as a
1315 * single-wire arith PLUS REFERENCING the survivor, so a shared
1316 * survivor RV keeps its single gate identity (and single MC draw);
1317 * foldSemiringIdentities then collapses the passthrough wrapper, and
1318 * a constant survivor lets runConstantFold reduce the enclosing sum
1319 * (e.g. avg's provenance-weighted count) to a gate_value. */
1320 if (pi == 1.0) {
1321 gc.liftConditionedToTarget(g, wires[1]);
1322 ++counter;
1323 } else if (pi == 0.0) {
1324 gc.liftConditionedToTarget(g, wires[2]);
1325 ++counter;
1326 }
1327 }
1328 return counter;
1329}
1330
1332{
1333 unsigned counter = 0;
1334
1335 /* Base RVs that couple two or more sibling subtrees must not be folded
1336 * into a fresh (independent) identity. Census once over the loaded
1337 * circuit. A leaf is coupling-critical when it is reachable (through
1338 * arith / latent-parameter edges) from two or more sibling contexts that
1339 * are combined NON-additively -- i.e. either
1340 * (A) two or more wires of a gate_cmp [the two sides of a cmp], or
1341 * (B) two or more arms of a non-additive gate_arith combinator
1342 * (MIN / MAX / PERCENTILE order statistics, MINUS / TIMES / DIV /
1343 * POW), e.g. the shared T0 in min(T0+T1, T0+T2).
1344 * Folding each such sibling arm into a fresh gate_rv would orphan the
1345 * shared leaf and silently apply the independence approximation to a
1346 * draw the order statistic / product / comparison must see jointly.
1347 * PLUS is deliberately excluded from (B): a leaf repeated across addends
1348 * (x + 2x) is consolidated by try_plus_aggregate while preserving its
1349 * identity, and try_sum_closure already bails on a within-sum repeat, so
1350 * an additive sum never decouples a shared leaf. Folds only ever remove
1351 * references to a leaf, so the load-time census is a sound conservative
1352 * guard for the whole run. The identity-minting folds consult
1353 * @c g_shared_base_rvs via @c is_shared_base_rv. */
1354 std::unordered_set<gate_t> shared_base_rvs;
1355 {
1356 const auto nb = gc.getNbGates();
1357 /* (A) sharing across the sides of a comparator. */
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);
1361 if (gc.getGateType(g) != gate_cmp) continue;
1362 for (gate_t w : gc.getWires(g)) {
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];
1366 }
1367 }
1368 for (const auto &[rv, n] : cmp_footprint_count)
1369 if (n > 1) shared_base_rvs.insert(rv);
1370 /* (B) sharing across the arms of a non-additive arith combinator. */
1371 for (std::size_t i = 0; i < nb; ++i) {
1372 auto g = static_cast<gate_t>(i);
1373 if (gc.getGateType(g) != gate_arith) continue;
1374 if (static_cast<provsql_arith_op>(gc.getInfos(g).first) == PROVSQL_ARITH_PLUS)
1375 continue;
1376 const auto &arms = gc.getWires(g);
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];
1383 }
1384 for (const auto &[rv, n] : arm_count)
1385 if (n > 1) shared_base_rvs.insert(rv);
1386 }
1387 }
1388 g_shared_base_rvs = &shared_base_rvs;
1389 struct SharedGuard {
1390 ~SharedGuard() { g_shared_base_rvs = nullptr; }
1391 } shared_guard;
1392
1393 /* Pass 1: bottom-up DFS applying every rule EXCEPT the scalar-times-RV
1394 * fold. Deferring that one rule lets @c try_plus_aggregate see
1395 * @c arith(TIMES, value:c, X) shapes inside a parent PLUS -- the
1396 * decomposer recognises them as @c c·X with @c rv_gate=X, so a
1397 * sibling @c x in @c 2·x + x correctly aggregates to coefficient
1398 * three on the shared base RV. If the scalar fold had fired bottom-up
1399 * on the inner TIMES first it would have minted a fresh @c gate_rv
1400 * there, decoupling its identity from the sibling @c x and forcing
1401 * the outer sum-closure path which assumes independence. */
1402 {
1403 std::unordered_set<gate_t> done;
1404 const auto nb = gc.getNbGates();
1405 for (std::size_t i = 0; i < nb; ++i) {
1406 simplify(gc, static_cast<gate_t>(i), done, counter,
1407 /*include_scalar_fold=*/false);
1408 }
1409 }
1410
1411 /* Pass 2: scalar-times-RV fold and NEG-of-RV fold on every
1412 * remaining @c gate_arith. Pass 1's aggregator and family closures
1413 * have already consumed the shapes where these folds would have
1414 * lost shared-RV identity; any surviving 2-wire
1415 * <tt>arith(TIMES, value:c, gate_rv)</tt> or 1-wire
1416 * <tt>arith(NEG, gate_rv)</tt> is now either standalone (no sibling
1417 * to couple with) or the leftover wrapper from a single-RV
1418 * aggregation result. No DFS is needed -- the rules are local and
1419 * idempotent, and walking the gate range with the post-pass-1
1420 * @c getNbGates() picks up the freshly minted wrappers from
1421 * @c try_plus_aggregate, @c try_mixture_lift, and the
1422 * MINUS-to-PLUS canonicalisation. */
1423 {
1424 const auto nb = gc.getNbGates();
1425 for (std::size_t i = 0; i < nb; ++i) {
1426 auto g = static_cast<gate_t>(i);
1427 if (gc.getGateType(g) == gate_arith) {
1428 if (try_times_scalar_rv(gc, g)) ++counter;
1429 else if (try_neg_rv(gc, g)) ++counter;
1430 }
1431 }
1432 }
1433
1434 return counter;
1435}
1436
1437namespace {
1438
1439/**
1440 * @brief Test whether both sides of @p cmp_gate are a continuous-only
1441 * island (subtree of @c gate_value / @c gate_rv / @c gate_arith).
1442 *
1443 * A continuous island has no Boolean / aggregate / IO gates underneath
1444 * the cmp; the only outward edge is the cmp itself. This is the
1445 * shape monteCarloRV's @c evalScalar can integrate over, so per-cmp
1446 * MC marginalisation is sound on these and these alone.
1447 */
1448bool is_continuous_island_cmp(const GenericCircuit &gc, gate_t cmp_gate)
1449{
1450 const auto &wires = gc.getWires(cmp_gate);
1451 if (wires.size() != 2) return false;
1452
1453 std::unordered_set<gate_t> seen;
1454 std::stack<gate_t> stk;
1455 stk.push(wires[0]);
1456 stk.push(wires[1]);
1457 while (!stk.empty()) {
1458 gate_t g = stk.top(); stk.pop();
1459 if (!seen.insert(g).second) continue;
1460 auto t = gc.getGateType(g);
1461 if (t == gate_value || t == gate_rv || t == gate_arith) {
1462 for (gate_t c : gc.getWires(g)) stk.push(c);
1463 continue;
1464 }
1465 if (t == gate_mixture) {
1466 /* Categorical-form mixture (from @c provsql.categorical): a
1467 * discrete scalar leaf with no continuous identities below.
1468 * Treat it as a black-box scalar leaf and don't descend. */
1469 if (gc.isCategoricalMixture(g)) continue;
1470 /* Classic 3-wire mixture: first wire is a gate_input Bernoulli;
1471 * the rest of the island walker would reject it as
1472 * non-continuous, but the Monte-Carlo sampler handles it
1473 * correctly via per-iteration coupling. Treat the mixture as
1474 * a black-box scalar leaf in the island shape: do NOT descend
1475 * into wires[0], only into the scalar branches wires[1] /
1476 * wires[2]. */
1477 const auto &mw = gc.getWires(g);
1478 if (mw.size() != 3) return false;
1479 stk.push(mw[1]);
1480 stk.push(mw[2]);
1481 continue;
1482 }
1483 return false;
1484 }
1485 return true;
1486}
1487
1488/**
1489 * @brief Collect the base @c gate_rv leaves reachable from @p root
1490 * through @c gate_arith composition.
1491 *
1492 * The set is the cmp's "RV footprint": two cmps share an island iff
1493 * their footprints overlap (a shared base RV is the only way their
1494 * sampled values can be correlated, given the island shape).
1495 */
1496void collect_cmp_rv_footprint(const GenericCircuit &gc, gate_t cmp_gate,
1497 std::unordered_set<gate_t> &fp)
1498{
1499 std::unordered_set<gate_t> seen;
1500 std::stack<gate_t> stk;
1501 for (gate_t w : gc.getWires(cmp_gate)) stk.push(w);
1502 while (!stk.empty()) {
1503 gate_t g = stk.top(); stk.pop();
1504 if (!seen.insert(g).second) continue;
1505 auto t = gc.getGateType(g);
1506 if (t == gate_rv) {
1507 fp.insert(g);
1508 /* A compound/latent RV carries its distribution parameters as
1509 * wires (e.g. normal($0, 1) with $0 a shared scalar leaf). Two
1510 * comparators over RVs that share a latent parameter are
1511 * correlated through it, so descend into the parameter subtrees
1512 * and collect their base RVs as well. A bare RV has no wires and
1513 * this is a no-op. */
1514 for (gate_t c : gc.getWires(g)) stk.push(c);
1515 continue;
1516 }
1517 if (t == gate_arith) {
1518 for (gate_t c : gc.getWires(g)) stk.push(c);
1519 continue;
1520 }
1521 if (t == gate_mixture) {
1522 /* Categorical-form mixture (from @c provsql.categorical):
1523 * discrete leaves, no continuous identities below. Stop. */
1524 if (gc.isCategoricalMixture(g)) continue;
1525 /* Classic 3-wire mixture: descend into the scalar branches but
1526 * NOT into the Bernoulli (wires[0] is a gate_input, not a
1527 * continuous RV identity). Two cmps that share a mixture's
1528 * continuous RVs still need to be grouped together; sharing the
1529 * Bernoulli alone does too, but that coupling is captured at
1530 * the sampler level rather than here -- the joint-table sampler
1531 * hits both cmps in the same MC iteration and the shared
1532 * bool_cache_ produces coherent draws. */
1533 const auto &mw = gc.getWires(g);
1534 if (mw.size() == 3) { stk.push(mw[1]); stk.push(mw[2]); }
1535 continue;
1536 }
1537 /* gate_value contributes no RV identity; other types should not
1538 * appear here (is_continuous_island_cmp gates that path), but if
1539 * they did we'd simply ignore them in the footprint &ndash; the
1540 * decomposer's safety relies on the island-shape pre-check, not
1541 * on this routine. */
1542 }
1543}
1544
1545/**
1546 * @brief Collect the classic 3-wire @c gate_mixture selector wires
1547 * (the Bernoulli @c wires[0]) reachable from @p cmp_gate through
1548 * @c gate_arith / @c gate_mixture composition.
1549 *
1550 * Mirrors @c collect_cmp_rv_footprint's walk but records each mixture's
1551 * Boolean SELECTOR rather than its continuous base RVs, descending into
1552 * the value arms so nested mixtures contribute their selectors too. A
1553 * mixture's selector is a latent Boolean whose sampled value the mixture
1554 * consumes; if it is shared with the rest of the circuit (another
1555 * mixture, or an external Boolean use such as conditioning), the cmp's
1556 * truth value is coupled to those sites. Marginalising the cmp into an
1557 * independent Bernoulli would then decorrelate the selector from its
1558 * other uses &ndash; a semantics change &ndash; so the decomposer must
1559 * leave such a cmp for the whole-circuit MC sampler (which couples the
1560 * selector across all its uses via the per-iteration @c bool_cache_).
1561 * Categorical-form mixtures carry no single Bernoulli selector and are
1562 * skipped. */
1563void collect_cmp_mixture_selectors(const GenericCircuit &gc, gate_t cmp_gate,
1564 std::unordered_set<gate_t> &sels)
1565{
1566 std::unordered_set<gate_t> seen;
1567 std::stack<gate_t> stk;
1568 for (gate_t w : gc.getWires(cmp_gate)) stk.push(w);
1569 while (!stk.empty()) {
1570 gate_t g = stk.top(); stk.pop();
1571 if (!seen.insert(g).second) continue;
1572 auto t = gc.getGateType(g);
1573 if (t == gate_arith) {
1574 for (gate_t c : gc.getWires(g)) stk.push(c);
1575 continue;
1576 }
1577 if (t == gate_mixture) {
1578 if (gc.isCategoricalMixture(g)) continue;
1579 const auto &mw = gc.getWires(g);
1580 if (mw.size() == 3) {
1581 sels.insert(mw[0]);
1582 stk.push(mw[1]);
1583 stk.push(mw[2]);
1584 }
1585 continue;
1586 }
1587 /* gate_rv / gate_value: no selector below. */
1588 }
1589}
1590
1591} // namespace
1592
1593namespace {
1594
1595/* Joint-table cap. 2^k mulinput leaves are materialised per group;
1596 * 256 cells is more than ample for HAVING/WHERE workloads while
1597 * keeping the in-memory footprint and the per-cell MC variance
1598 * (samples / 2^k counts per cell) bounded. Groups exceeding the
1599 * cap fall through to whole-circuit MC by leaving their cmps as
1600 * gate_cmp; the dispatch in probability_evaluate then routes
1601 * through monteCarloRV. */
1602constexpr std::size_t JOINT_TABLE_K_MAX = 8;
1603
1604/**
1605 * @brief Test whether @c AnalyticEvaluator would resolve @p cmp_gate
1606 * analytically on its own.
1607 *
1608 * The decomposer now runs before @c AnalyticEvaluator (so shared
1609 * bare-RV cmps reach the grouping logic and the fast path's
1610 * analytical CDF can fire), but it must leave isolated bare-RV cmps
1611 * untouched: marginalising those via MC would waste samples on a
1612 * case the closed-form CDF handles exactly. Mirror the shape match
1613 * in @c tryAnalyticDecide (bare RV vs gate_value either way around;
1614 * two bare normal RVs).
1615 */
1616bool is_analytic_singleton_cmp(const GenericCircuit &gc, gate_t cmp_gate)
1617{
1618 const auto &wires = gc.getWires(cmp_gate);
1619 if (wires.size() != 2) return false;
1620 auto t0 = gc.getGateType(wires[0]);
1621 auto t1 = gc.getGateType(wires[1]);
1622
1623 /* X cmp c / c cmp X: AnalyticEvaluator resolves any supported
1624 * distribution kind via the closed-form CDF. */
1625 if ((t0 == gate_rv && t1 == gate_value) ||
1626 (t0 == gate_value && t1 == gate_rv))
1627 return true;
1628
1629 /* Categorical-form mixture cmp constant: AnalyticEvaluator's
1630 * @c categoricalDecide computes the exact mass sum over the
1631 * mulinputs satisfying the predicate, so the decomposer should not
1632 * pre-empt with per-cmp MC. Also picks up the
1633 * @c try_categorical_mixture_lift output (a constant scaled / offset
1634 * categorical), keeping the analytical path end-to-end for
1635 * <tt>c · X cmp k</tt> shapes over categorical RVs. */
1636 if ((gc.isCategoricalMixture(wires[0]) && t1 == gate_value) ||
1637 (gc.isCategoricalMixture(wires[1]) && t0 == gate_value))
1638 return true;
1639
1640 /* X cmp Y, two distinct bare RVs: AnalyticEvaluator's @c rvVsRvDecide
1641 * decides it -- a same-family closed form (Normal-Normal, Exp-Exp,
1642 * Uniform-Uniform) or the mixed-family 1-D quadrature. Two distinct
1643 * bare-RV leaves are independent, so leaving them for that path is exact
1644 * (or high-accuracy), never a per-cmp MC. */
1645 if (t0 == gate_rv && t1 == gate_rv) {
1646 auto sx = parse_distribution_spec(gc.getExtra(wires[0]));
1647 auto sy = parse_distribution_spec(gc.getExtra(wires[1]));
1648 if (sx && sy)
1649 return true;
1650 }
1651 return false;
1652}
1653
1654/**
1655 * @brief Information needed by @c inline_fast_path: the shared scalar
1656 * plus, for each cmp, the comparison operator and the
1657 * constant rhs threshold (after flipping for cmps shaped
1658 * @c c @c op @c X).
1659 */
1660struct FastPathInfo {
1661 gate_t scalar;
1662 std::vector<ComparisonOperator> ops; /* one per cmp, oriented as `scalar op c` */
1663 std::vector<double> thresholds; /* one per cmp */
1664};
1665
1667{
1668 switch (op) {
1675 }
1676 return op;
1677}
1678
1679bool apply_cmp(double l, ComparisonOperator op, double r)
1680{
1681 switch (op) {
1682 case ComparisonOperator::LT: return l < r;
1683 case ComparisonOperator::LE: return l <= r;
1684 case ComparisonOperator::EQ: return l == r;
1685 case ComparisonOperator::NE: return l != r;
1686 case ComparisonOperator::GE: return l >= r;
1687 case ComparisonOperator::GT: return l > r;
1688 }
1689 return false;
1690}
1691
1692/**
1693 * @brief Detect the monotone-shared-scalar fast path on a group of
1694 * comparators.
1695 *
1696 * Fires when every cmp in @p cmps has one side equal to a single
1697 * shared gate_t @c s and the other side a @c gate_value: the k cmps
1698 * then jointly partition the @c s-line into at most k+1 intervals,
1699 * with each interval producing a deterministic k-bit outcome. This
1700 * shape is common in HAVING / WHERE with multiple thresholds on the
1701 * same aggregate / column: e.g.
1702 * <tt>count(*) > 10 OR count(*) < 5</tt>.
1703 *
1704 * Returns @c std::nullopt when any cmp has both non-constant sides,
1705 * when the cmps don't all share the same @c s gate_t, when a
1706 * comparator OID is unrecognised, or when @c EQ / @c NE appears (the
1707 * interval representation can't express a measure-zero point).
1708 */
1709std::optional<FastPathInfo>
1710detect_shared_scalar(const GenericCircuit &gc,
1711 const std::vector<gate_t> &cmps)
1712{
1713 FastPathInfo info;
1714 info.ops.reserve(cmps.size());
1715 info.thresholds.reserve(cmps.size());
1716 bool first = true;
1717
1718 for (gate_t c : cmps) {
1719 const auto &wires = gc.getWires(c);
1720 if (wires.size() != 2) return std::nullopt;
1721
1722 bool ok = false;
1723 ComparisonOperator op = cmpOpFromOid(gc.getInfos(c).first, ok);
1724 if (!ok) return std::nullopt;
1725 /* EQ / NE on continuous RVs have measure zero / one and were
1726 * already resolved by RangeCheck; if we still see one we don't
1727 * know how to fit it into an interval partition. Bail. */
1729 return std::nullopt;
1730
1731 gate_t scalar_side = static_cast<gate_t>(-1);
1732 double threshold = std::numeric_limits<double>::quiet_NaN();
1733 ComparisonOperator effective_op = op;
1734 if (gc.getGateType(wires[1]) == gate_value) {
1735 scalar_side = wires[0];
1736 try { threshold = parseDoubleStrict(gc.getExtra(wires[1])); }
1737 catch (const CircuitException &) { return std::nullopt; }
1738 } else if (gc.getGateType(wires[0]) == gate_value) {
1739 scalar_side = wires[1];
1740 try { threshold = parseDoubleStrict(gc.getExtra(wires[0])); }
1741 catch (const CircuitException &) { return std::nullopt; }
1742 effective_op = flip_cmp_op(op);
1743 } else {
1744 return std::nullopt;
1745 }
1746
1747 if (first) {
1748 info.scalar = scalar_side;
1749 first = false;
1750 } else if (info.scalar != scalar_side) {
1751 return std::nullopt;
1752 }
1753 info.ops.push_back(effective_op);
1754 info.thresholds.push_back(threshold);
1755 }
1756 return info;
1757}
1758
1759/**
1760 * @brief Inline a fast-path joint table for a monotone-shared-scalar
1761 * group.
1762 *
1763 * The k cmps partition the scalar line into at most k+1 intervals
1764 * (one per pair of consecutive sorted distinct thresholds plus the
1765 * two infinite tails). Each interval gets a single mulinput with
1766 * probability equal to the scalar's mass on the interval; the
1767 * comparator outcomes are deterministic per interval (evaluated at
1768 * a strictly-interior representative point) and the k cmps are
1769 * rewritten as @c gate_plus over the mulinputs whose interval makes
1770 * them true.
1771 *
1772 * Interval probabilities are computed analytically via @c cdfAt when
1773 * the scalar is a bare @c gate_rv with a CDF the helper supports;
1774 * otherwise (a @c gate_arith composite, or an Erlang with
1775 * non-integer shape) we fall back to MC by sampling the scalar
1776 * @p samples times and binning into intervals.
1777 *
1778 * Returns @c true when the group was resolved. When the analytical
1779 * CDF is unavailable and @p allow_mc is false (@c rv_mc_samples = 0),
1780 * returns @c false without touching the circuit: the caller then
1781 * raises rather than letting each cmp collapse to an independent
1782 * marginal (which would silently return the product of the marginals
1783 * for correlated events).
1784 */
1785bool inline_fast_path(GenericCircuit &gc,
1786 const std::vector<gate_t> &cmps,
1787 const FastPathInfo &info,
1788 unsigned samples,
1789 bool allow_mc)
1790{
1791 /* Sort + dedup thresholds; the resulting m distinct boundaries
1792 * partition R into m+1 open intervals
1793 * (-∞, t_0), (t_0, t_1), ..., (t_{m-1}, +∞). */
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;
1799
1800 /* Compute interval probabilities. Try the analytical CDF first:
1801 * when the shared scalar is a bare @c gate_rv with a CDF
1802 * @c cdfAt understands, the interval probability is
1803 * @c F(t_{i+1}) - F(t_i) exactly &ndash; no MC noise, no sampling.
1804 * This is the headline benefit of the fast path: shared bare-RV
1805 * groups land on the exact dependent truth and the resulting
1806 * Bernoulli probabilities propagate through tree-decomposition /
1807 * compilation without any sampling noise contributed by the
1808 * decomposer. Fall back to MC binning over @p samples scalar
1809 * draws when the scalar is a @c gate_arith composite (no CDF) or
1810 * when @c cdfAt returns NaN on a boundary (Erlang with
1811 * non-integer shape, etc.). */
1812 std::vector<double> interval_probs(nb_intervals, 0.0);
1813 bool analytical = false;
1814 if (gc.getGateType(info.scalar) == gate_rv) {
1815 auto spec = parse_distribution_spec(gc.getExtra(info.scalar));
1816 if (spec) {
1817 std::vector<double> cdf_at_boundary(m);
1818 bool all_ok = true;
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; }
1822 }
1823 if (all_ok) {
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];
1828 analytical = true;
1829 }
1830 }
1831 }
1832 if (!analytical) {
1833 /* No closed-form CDF for the shared scalar (a gate_arith composite,
1834 * or a non-integer-shape Erlang): the joint needs MC binning. With
1835 * MC disabled we cannot resolve it correctly -- decline so the caller
1836 * raises, rather than leaving the cmps for an independent per-cmp
1837 * collapse that would silently return the product of the marginals. */
1838 if (!allow_mc) return false;
1839 auto draws = monteCarloScalarSamples(gc, info.scalar, samples);
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];
1844 }
1845 for (auto &p : interval_probs) p /= samples;
1846 }
1847
1848 /* For each interval, determine the k-bit cmp outcome word. Pick
1849 * a representative point strictly inside the interval: the
1850 * midpoint for finite intervals, t_0 - 1 / t_{m-1} + 1 for the
1851 * infinite tails. Continuous distributions assign zero mass to
1852 * the boundaries, so the choice of interior point doesn't
1853 * affect any cmp's outcome on the open interval. */
1854 std::vector<unsigned long> outcome_word(nb_intervals, 0);
1855 for (std::size_t i = 0; i < nb_intervals; ++i) {
1856 double point;
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]))
1863 w |= (1ul << j);
1864 }
1865 outcome_word[i] = w;
1866 }
1867
1868 /* Allocate key + per-interval mulinputs (skipping zero-prob
1869 * intervals to keep the materialised circuit lean). */
1870 gate_t key = gc.addAnonymousInputGate(1.0);
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] =
1876 gc.addAnonymousMulinputGate(key, interval_probs[i],
1877 static_cast<unsigned>(i));
1878 }
1879
1880 /* Rewrite each cmp as gate_plus over the mulinputs whose
1881 * interval-outcome word has the cmp's bit set. */
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);
1890 }
1891 gc.resolveToPlus(cmps[j], std::move(plus_wires));
1892 }
1893 return true;
1894}
1895
1896/**
1897 * @brief Inline a joint-distribution table over a group of k cmps
1898 * sharing an island.
1899 *
1900 * Materialises 2^k - z mulinput leaves (where z is the number of
1901 * outcomes with empirical probability zero, omitted to keep the
1902 * circuit lean), all sharing a fresh anonymous key gate. Each
1903 * comparator @c cmps[i] is rewritten in place as @c gate_plus over
1904 * the mulinputs whose joint outcome word has bit @c i set; the
1905 * combined probability is the marginal P(cmp_i = 1) and shared bits
1906 * across different cmps reuse the same mulinput leaf so the OR over
1907 * cmps at downstream sites correctly observes the joint distribution
1908 * (mutually exclusive over the joint outcomes).
1909 *
1910 * Sound when the per-iteration sampler memoisation in
1911 * @c monteCarloRV / @c monteCarloJointDistribution gives all k cmps
1912 * a consistent draw of the shared island - which is precisely the
1913 * is_continuous_island_cmp + shared-footprint precondition the
1914 * caller has already enforced.
1915 */
1916void inline_joint_table(GenericCircuit &gc,
1917 const std::vector<gate_t> &cmps,
1918 unsigned samples)
1919{
1920 const unsigned k = static_cast<unsigned>(cmps.size());
1921 auto probs = monteCarloJointDistribution(gc, cmps, samples);
1922
1923 /* Fresh key gate (the anonymous block anchor for these mulinputs).
1924 * Probability 1.0 because the key itself is not a sampled choice;
1925 * the mutually-exclusive outcomes among the mulinputs are what
1926 * carries the joint mass. */
1927 gate_t key = gc.addAnonymousInputGate(1.0);
1928
1929 /* Allocate one mulinput per joint outcome with positive probability.
1930 * Zero-probability outcomes are pruned: the cmp gate_plus
1931 * rewrites below would have included them as wires with prob 0,
1932 * which is a no-op in OR (gate_zero is the additive identity).
1933 * value_index = w gives independentEvaluation's mulin_seen dedup
1934 * a stable key (group, info) per outcome. */
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] =
1941 gc.addAnonymousMulinputGate(key, probs[w],
1942 static_cast<unsigned>(w));
1943 }
1944
1945 /* Rewrite each cmp as gate_plus over the mulinputs whose joint
1946 * outcome word has the cmp's bit set. */
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);
1955 }
1956 gc.resolveToPlus(cmps[i], std::move(plus_wires));
1957 }
1958}
1959
1960/**
1961 * @brief A group of comparisons all sharing one pivot bare RV X, each against
1962 * an independent bare RV or a constant. Unlike the monotone-shared-
1963 * scalar fast path (all thresholds constant on one line), the other
1964 * operands are themselves random, so the joint of the k comparisons is
1965 * a 2^k table of pivot-conjunction integrals rather than a partition of
1966 * one line into intervals. This is the RV-vs-RV correlated-join island,
1967 * e.g. @c "(x>y) AND (x>z)" or the conditioning @c "(x>y) | (x>z)".
1968 */
1969struct PivotIslandInfo {
1970 DistributionSpec pivotSpec;
1971 struct Factor {
1972 bool isConst;
1973 DistributionSpec other; /* valid iff !isConst */
1974 double konst; /* valid iff isConst */
1975 bool trueIsGreater; /* cmp true <=> X > operand */
1976 };
1977 std::vector<Factor> factors; /* one per cmp, index = bit position */
1978};
1979
1980/* Detect a shared-pivot-RV island: every cmp compares a common pivot bare RV X
1981 * against an independent bare RV (distinct leaf, appearing once) or a constant.
1982 * nullopt otherwise (no common pivot, a shared/repeated other RV, an agg cmp,
1983 * EQ/NE, or a non-bare operand). */
1984std::optional<PivotIslandInfo>
1985detect_shared_pivot_rv(const GenericCircuit &gc,
1986 const std::vector<gate_t> &cmps)
1987{
1988 gate_t pivot = static_cast<gate_t>(-1);
1989 bool havePivot = false;
1990 std::optional<DistributionSpec> pivotSpec;
1991 PivotIslandInfo info;
1992 std::unordered_set<gate_t> othersSeen;
1993
1994 for (gate_t c : cmps) {
1995 if (gc.getGateType(c) != gate_cmp) return std::nullopt;
1996 const auto &w = gc.getWires(c);
1997 if (w.size() != 2) return std::nullopt;
1998 bool ok = false;
1999 ComparisonOperator op = cmpOpFromOid(gc.getInfos(c).first, ok);
2000 if (!ok || op == ComparisonOperator::EQ || op == ComparisonOperator::NE)
2001 return std::nullopt;
2002
2003 /* Find which side is the pivot: a bare gate_rv consistent across cmps. */
2004 gate_t pv, other; bool pivotLeft;
2005 if (gc.getGateType(w[0]) == gate_rv &&
2006 (!havePivot || w[0] == pivot)) { pv = w[0]; other = w[1]; pivotLeft = true; }
2007 else if (gc.getGateType(w[1]) == gate_rv &&
2008 (!havePivot || w[1] == pivot)) { pv = w[1]; other = w[0]; pivotLeft = false; }
2009 else return std::nullopt;
2010
2011 if (!havePivot) {
2012 auto sp = parse_distribution_spec(gc.getExtra(pv));
2013 if (!sp) return std::nullopt;
2014 pivot = pv; pivotSpec = *sp; havePivot = true;
2015 }
2016
2017 const bool greaterOp = (op == ComparisonOperator::GT ||
2019 const bool lessOp = (op == ComparisonOperator::LT ||
2021 const bool trueIsGreater = pivotLeft ? greaterOp : lessOp;
2022
2023 PivotIslandInfo::Factor f;
2024 f.trueIsGreater = trueIsGreater;
2025 if (gc.getGateType(other) == gate_value) {
2026 f.isConst = true;
2027 try { f.konst = parseDoubleStrict(gc.getExtra(other)); }
2028 catch (const CircuitException &) { return std::nullopt; }
2029 } else if (gc.getGateType(other) == gate_rv && other != pivot) {
2030 if (!othersSeen.insert(other).second) return std::nullopt; /* shared -> dependent */
2031 auto sp = parse_distribution_spec(gc.getExtra(other));
2032 if (!sp) return std::nullopt;
2033 f.isConst = false;
2034 f.other = *sp;
2035 } else return std::nullopt;
2036 info.factors.push_back(std::move(f));
2037 }
2038 if (!havePivot) return std::nullopt;
2039 info.pivotSpec = *pivotSpec;
2040 return info;
2041}
2042
2043/* Install an analytic 2^k joint table for a shared-pivot-RV island. Cell
2044 * probability for outcome word w is P(∧_j cmp_j == bit_j(w)) computed as the
2045 * pivot-conjunction integral ∫ f_X(x) Π_j W_j(x) dx, where W_j weights the
2046 * region making cmp_j equal to its bit: an RV factor contributes F_Y(x) or
2047 * 1-F_Y(x), a constant factor clips the integration window. Mirrors
2048 * @c inline_joint_table's circuit surgery (one key, one mulinput per positive
2049 * outcome, each cmp -> gate_plus over the outcomes with its bit set), so the
2050 * downstream Boolean OR/AND observes the correct correlated joint. Returns
2051 * false (touching nothing) if a density/CDF is undefined, so the caller can
2052 * raise or fall back to MC. */
2053bool inline_analytic_pivot_joint_table(GenericCircuit &gc,
2054 const std::vector<gate_t> &cmps,
2055 const PivotIslandInfo &info)
2056{
2057 const unsigned k = static_cast<unsigned>(cmps.size());
2058 const std::size_t nb_outcomes = std::size_t{1} << k;
2059
2060 const auto dX = makeDistribution(info.pivotSpec);
2061 double lo0, hi0;
2062 if (!dX->integrationRange(lo0, hi0)) return false;
2063
2064 /* Construct each RV factor's distribution once. */
2065 std::vector<std::unique_ptr<Distribution>> otherDist(k);
2066 for (unsigned j = 0; j < k; ++j)
2067 if (!info.factors[j].isConst)
2068 otherDist[j] = makeDistribution(info.factors[j].other);
2069
2070 std::vector<double> probs(nb_outcomes, 0.0);
2071 for (std::size_t w = 0; w < nb_outcomes; ++w) {
2072 /* Constant factors clip the window for this cell. */
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;
2077 /* "X > c" holds in this cell iff (trueIsGreater == bit). */
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);
2081 }
2082 if (!(hi > lo)) { probs[w] = 0.0; continue; }
2083
2084 const double cell = simpsonIntegrate(lo, hi, kSimpsonPanels,
2085 [&](double x) {
2086 const double fX = dX->pdf(x);
2087 if (std::isnan(fX)) return std::numeric_limits<double>::quiet_NaN();
2088 double weight = fX;
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);
2096 }
2097 return weight;
2098 });
2099 if (std::isnan(cell)) return false;
2100 probs[w] = cell;
2101 }
2102
2103 /* Circuit surgery: one shared key, one mulinput per positive outcome, each
2104 * cmp rewritten as gate_plus over the outcomes whose bit it owns. */
2105 gate_t key = gc.addAnonymousInputGate(1.0);
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] =
2110 gc.addAnonymousMulinputGate(key, probs[w], static_cast<unsigned>(w));
2111 }
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);
2119 }
2120 gc.resolveToPlus(cmps[i], std::move(plus_wires));
2121 }
2122 return true;
2123}
2124
2125} // namespace
2126
2127unsigned runHybridDecomposer(GenericCircuit &gc, unsigned samples)
2128{
2129 /* @c rv_mc_samples = 0 does NOT short-circuit the decomposer: the
2130 * monotone-shared-scalar fast path resolves a group of comparisons
2131 * against constants on one bare @c gate_rv analytically (via the CDF,
2132 * no sampling), which is exactly the correlation-aware joint a
2133 * conditioning like @c "(x >= 2000) | (x >= 1000)" needs. Skipping
2134 * the decomposer here would leave those cmps for @c runAnalyticEvaluator
2135 * to collapse one at a time, silently returning the product of the
2136 * marginals for correlated events. Only the genuinely MC-bound arms
2137 * (a composite shared scalar, an RV-vs-RV joint table, a non-analytic
2138 * singleton) are gated on @p allow_mc; under @c samples = 0 a correlated
2139 * island with no closed form raises rather than falling back silently. */
2140 const bool allow_mc = (samples > 0);
2141
2142 /* Snapshot all gate_cmp ids that look like continuous islands.
2143 * Each call later mutates a snapshot entry from @c gate_cmp to
2144 * @c gate_input via @c resolveCmpToBernoulli (singleton group)
2145 * or to @c gate_plus via @c resolveToPlus (multi-cmp group), but
2146 * the snapshot vector is unaffected. The defensive type re-check
2147 * at iteration time guards against intervening mutations. */
2148 const auto nb = gc.getNbGates();
2149 std::vector<gate_t> cmps;
2150 for (std::size_t i = 0; i < nb; ++i) {
2151 auto g = static_cast<gate_t>(i);
2152 if (gc.getGateType(g) == gate_cmp && is_continuous_island_cmp(gc, g))
2153 cmps.push_back(g);
2154 }
2155
2156 /* Compute the per-cmp footprint up front so the pairwise-overlap
2157 * check is O(C * C * F) rather than O(C * C * tree_size). */
2158 std::unordered_map<gate_t, std::unordered_set<gate_t>> footprints;
2159 footprints.reserve(cmps.size());
2160 for (gate_t c : cmps) {
2161 collect_cmp_rv_footprint(gc, c, footprints[c]);
2162 }
2163
2164 /* Wire in-degree census over the whole circuit: indeg[g] is the number
2165 * of parents that reference g anywhere. Compared per group against the
2166 * references seen WITHIN the group's own island, it tells whether a
2167 * mixture selector's Boolean sub-DAG is observed outside the group (a
2168 * coupling the island-local marginalisation cannot preserve). */
2169 std::unordered_map<gate_t, unsigned> indeg;
2170 {
2171 for (std::size_t i = 0; i < nb; ++i)
2172 for (gate_t c : gc.getWires(static_cast<gate_t>(i)))
2173 ++indeg[c];
2174 }
2175
2176 /* Group cmps into connected components by base-RV footprint
2177 * overlap (union-find via parent[]). Linear-probe path
2178 * compression keeps the asymptotics near-linear in the number of
2179 * pairwise overlap checks. */
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]];
2185 x = parent[x];
2186 }
2187 return x;
2188 };
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;
2192 };
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; }
2202 }
2203 }
2204 }
2205
2206 /* Collect cmps by component root. */
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]);
2210
2211 unsigned resolved = 0;
2212 for (auto &[root, group] : groups) {
2213 (void) root;
2214 /* Defensive: re-check every cmp is still gate_cmp. Nothing in
2215 * the pipeline should have mutated them since the snapshot, but
2216 * the check is cheap. */
2217 bool all_pristine = true;
2218 for (gate_t c : group) {
2219 if (gc.getGateType(c) != gate_cmp) { all_pristine = false; break; }
2220 }
2221 if (!all_pristine) continue;
2222
2223 /* Semantics guard: an island-local resolution (a singleton marginal
2224 * Bernoulli, or a 2^k joint table over the group) replaces the cmps
2225 * with leaves that are independent of everything OUTSIDE this group.
2226 * That is sound only when the group's island is self-contained. The
2227 * channel the footprint grouping does not cover is a Bernoulli
2228 * MIXTURE selector: if a selector's Boolean sub-DAG is also observed
2229 * outside this island (another group's mixture, or an external
2230 * Boolean such as conditioning on the selector), marginalising here
2231 * would decorrelate it from those uses -- a semantics change. A
2232 * selector shared only WITHIN the group is fine: the same MC draw
2233 * that marginalises the group couples it internally. When a group
2234 * is not self-contained, leave all its cmps as raw gate_cmp for the
2235 * whole-circuit MC sampler, which couples every selector across all
2236 * its uses via the per-iteration bool_cache_. */
2237 {
2238 /* Island of this group: every gate reachable from its cmps, with a
2239 * count of how many references each gate receives from within it. */
2240 std::unordered_set<gate_t> island;
2241 std::unordered_map<gate_t, unsigned> island_ref;
2242 {
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); }
2249 }
2250 }
2251 /* A selector couples the group to the outside iff some gate in its
2252 * sub-DAG has a parent that is not part of this island (its global
2253 * in-degree exceeds the references seen within the island). */
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;
2263 for (gate_t w : gc.getWires(g)) st.push(w);
2264 }
2265 return false;
2266 };
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;
2270 for (gate_t s : sels)
2271 if (sub_dag_escapes_island(s)) { group_couples_selector = true; break; }
2272 if (group_couples_selector) continue;
2273 }
2274
2275 if (group.size() == 1) {
2276 /* Singleton island. If AnalyticEvaluator would resolve this
2277 * cmp exactly on its own (bare gate_rv vs gate_value, or two
2278 * bare normals), leave it untouched and let the closed-form
2279 * pass below handle it - no point burning MC samples on a
2280 * case with an analytical answer. Otherwise MC-marginalise
2281 * into a Bernoulli leaf here. */
2282 if (is_analytic_singleton_cmp(gc, group[0])) continue;
2283 /* A non-analytic singleton needs MC. With MC disabled leave the
2284 * cmp for the downstream "undecidable + rv_mc_samples = 0" raise
2285 * (a singleton has no cross-cmp correlation to lose). */
2286 if (!allow_mc) continue;
2287 double p = monteCarloRV(gc, group[0], samples);
2288 gc.resolveCmpToBernoulli(group[0], p);
2289 ++resolved;
2290 continue;
2291 }
2292
2293 /* Multi-cmp shared island. Try the monotone-shared-scalar fast
2294 * path first: when every cmp has shape `s op c` for a common
2295 * scalar gate_t s, the joint table is built from k+1 intervals
2296 * (analytical when s is a bare gate_rv with a known CDF, MC
2297 * binning otherwise) instead of 2^k cells, and the test
2298 * 14-style shared bare-RV case (`X > 0 OR X > 1`) lands on the
2299 * exact answer with no MC noise. When detection fails, fall
2300 * through to the generic 2^k MC joint table iff k is small
2301 * enough; larger groups keep their cmps as gate_cmp and fall
2302 * through to whole-circuit MC. */
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());
2306 continue;
2307 }
2308 /* Shared scalar with no closed-form CDF and MC disabled: raise
2309 * rather than let the cmps collapse to independent marginals. */
2310 throw CircuitException(
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)");
2315 }
2316
2317 /* Shared-pivot-RV island (e.g. `x>y AND x>z`, or the conditioning
2318 * `(x>y)|(x>z)`): the k comparisons share one pivot bare RV X against
2319 * independent operands, so the 2^k joint is a table of pivot-conjunction
2320 * integrals -- exact (quadrature), no MC. Preferred even when MC is
2321 * available (no sampling noise); resolves the `rv_mc_samples = 0` case
2322 * that would otherwise raise below. */
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());
2326 continue;
2327 }
2328 /* Integration declined (undefined density/CDF): fall through to MC or
2329 * the raise below. */
2330 }
2331
2332 /* Generic joint island (e.g. RV-vs-RV comparisons sharing a leaf):
2333 * only the 2^k MC joint table can evaluate it correctly. With MC
2334 * disabled, raise rather than leave the cmps for an independent
2335 * per-cmp collapse that silently returns the product of marginals. */
2336 if (!allow_mc)
2337 throw CircuitException(
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)");
2342
2343 if (group.size() > JOINT_TABLE_K_MAX) continue;
2344
2345 inline_joint_table(gc, group, samples);
2346 resolved += static_cast<unsigned>(group.size());
2347 }
2348
2349 return resolved;
2350}
2351
2352} // namespace provsql
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.
Definition Aggregation.h:39
@ LT
Less than (<).
Definition Aggregation.h:43
@ GT
Greater than (>).
Definition Aggregation.h:45
@ LE
Less than or equal (<=).
Definition Aggregation.h:42
@ NE
Not equal (<>).
Definition Aggregation.h:41
@ GE
Greater than or equal (>=).
Definition Aggregation.h:44
Closed-form CDF resolution for trivial gate_cmp shapes.
gate_t
Strongly-typed gate identifier.
Definition Circuit.h:49
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.
Definition Circuit.h:206
std::vector< gate_t > & getWires(gate_t g)
Return a mutable reference to the child-wire list of gate g.
Definition Circuit.h:140
gateType getGateType(gate_t g) const
Return the type of gate g.
Definition Circuit.h:130
std::vector< gate_t >::size_type getNbGates() const
Return the total number of gates in the circuit.
Definition Circuit.h:103
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)