52 std::vector<gate_t> &out)
62 if (!collectProductLeaves(gc, c, out))
return false;
70 std::vector<gate_t> &out)
73 if (!collectProductLeaves(gc, k, out))
77 std::sort(out.begin(), out.end());
78 if (std::adjacent_find(out.begin(), out.end()) != out.end())
103 const std::vector<unsigned> &ref)
105 std::map<gate_t, unsigned> internalRef;
106 std::set<gate_t> visited;
107 std::vector<gate_t> stk{agg};
109 while (!stk.empty()) {
110 gate_t g = stk.back(); stk.pop_back();
113 if (visited.insert(c).second) stk.push_back(c);
116 for (
gate_t g : visited) {
117 if (g == agg)
continue;
124 if (ref[
static_cast<std::size_t
>(g)] != internalRef[g])
131constexpr unsigned kMaxContributorLeaves = 20;
152 const std::vector<unsigned> &ref,
156 std::map<gate_t, unsigned> internalRef;
157 std::set<gate_t> seen;
158 std::vector<gate_t> order;
159 std::vector<std::pair<gate_t, bool>> stk{{g,
false}};
160 while (!stk.empty()) {
161 auto top = stk.back(); stk.pop_back();
167 if (top.second) { order.push_back(x);
continue; }
168 if (!seen.insert(x).second)
continue;
169 stk.push_back({x,
true});
173 stk.push_back({c,
false});
179 if (x == g)
continue;
184 if (ref[
static_cast<std::size_t
>(x)] != internalRef[x])
return false;
188 const int N =
static_cast<int>(order.size());
189 std::map<gate_t, int> pos;
190 for (
int i = 0; i < N; ++i) pos[order[i]] = i;
191 std::vector<gate_type> typ(N);
192 std::vector<std::vector<int>> childIdx(N);
193 std::vector<int> leafbit(N, -1);
194 std::vector<double> leafProb;
195 for (
int i = 0; i < N; ++i) {
198 leafbit[i] =
static_cast<int>(leafProb.size());
199 leafProb.push_back(gc.
getProb(order[i]));
201 for (
gate_t c : gc.
getWires(order[i])) childIdx[i].push_back(pos[c]);
204 const unsigned m =
static_cast<unsigned>(leafProb.size());
205 if (m > kMaxContributorLeaves)
return false;
208 std::vector<char> val(N);
210 for (uint32_t mask = 0; mask < (1u << m); ++mask) {
211 for (
int i = 0; i < N; ++i) {
215 case gate_input: val[i] = (mask >> leafbit[i]) & 1u;
break;
216 case gate_times: {
char v = 1;
for (
int c : childIdx[i]) v = v && val[c]; val[i] = v;
break; }
217 case gate_plus: {
char v = 0;
for (
int c : childIdx[i]) v = v || val[c]; val[i] = v;
break; }
218 case gate_monus: val[i] = val[childIdx[i][0]] && !val[childIdx[i][1]];
break;
219 default:
return false;
224 for (
unsigned b = 0; b < m; ++b)
225 pr *= (mask >> b) & 1u ? leafProb[b] : 1.0 - leafProb[b];
235 std::vector<int> parent;
236 explicit UnionFind(
int n) : parent(n) {
237 std::iota(parent.begin(), parent.end(), 0);
240 while (parent[x] != x) { parent[x] = parent[parent[x]]; x = parent[x]; }
243 void unite(
int a,
int b) { parent[find(a)] = find(b); }
250static std::vector<std::vector<int>> independenceBlocks(
251 const std::vector<std::vector<gate_t>> &contribs)
253 const std::size_t n = contribs.size();
254 UnionFind uf(
static_cast<int>(n));
255 std::map<gate_t, int> first_owner;
256 for (std::size_t i = 0; i < n; ++i)
257 for (
gate_t l : contribs[i]) {
258 auto it = first_owner.find(l);
259 if (it == first_owner.end()) first_owner[l] =
static_cast<int>(i);
260 else uf.unite(
static_cast<int>(i), it->second);
262 std::map<int, std::vector<int>> bmap;
263 for (std::size_t i = 0; i < n; ++i)
264 bmap[uf.find(
static_cast<int>(i))].push_back(
static_cast<int>(i));
265 std::vector<std::vector<int>> blocks;
266 blocks.reserve(bmap.size());
267 for (
auto &kv : bmap) blocks.push_back(std::move(kv.second));
274static std::vector<gate_t> commonLeaves(
275 const std::vector<std::vector<gate_t>> &contribs,
276 const std::vector<int> &members)
278 std::vector<gate_t> common = contribs[members[0]];
279 for (std::size_t mi = 1; mi < members.size() && !common.empty(); ++mi) {
280 std::vector<gate_t> tmp;
281 std::set_intersection(common.begin(), common.end(),
282 contribs[members[mi]].begin(),
283 contribs[members[mi]].end(),
284 std::back_inserter(tmp));
292static std::vector<std::vector<gate_t>> residualsOf(
293 const std::vector<std::vector<gate_t>> &contribs,
294 const std::vector<int> &members,
const std::vector<gate_t> &common)
296 std::vector<std::vector<gate_t>> residuals;
297 residuals.reserve(members.size());
298 for (
int m : members) {
299 std::vector<gate_t> r;
300 std::set_difference(contribs[m].begin(), contribs[m].end(),
301 common.begin(), common.end(), std::back_inserter(r));
302 residuals.push_back(std::move(r));
308static std::vector<double> convolve(
const std::vector<double> &a,
309 const std::vector<double> &b)
311 if (a.empty())
return b;
312 if (b.empty())
return a;
313 std::vector<double> r(a.size() + b.size() - 1, 0.0);
314 for (std::size_t i = 0; i < a.size(); ++i) {
315 if (a[i] == 0.0)
continue;
316 for (std::size_t j = 0; j < b.size(); ++j)
317 r[i + j] += a[i] * b[j];
325static std::vector<double> productConvolve(
const std::vector<double> &a,
326 const std::vector<double> &b)
328 if (a.empty() || b.empty())
return {};
329 const std::size_t amax = a.size() - 1, bmax = b.size() - 1;
330 std::vector<double> r(amax * bmax + 1, 0.0);
331 for (std::size_t i = 0; i <= amax; ++i) {
332 if (a[i] == 0.0)
continue;
333 for (std::size_t j = 0; j <= bmax; ++j)
334 r[i * j] += a[i] * b[j];
340struct ProductDecomp {
342 std::map<gate_t, int> leafFactor;
343 std::vector<std::vector<std::vector<gate_t>>> parts;
366static ProductDecomp decomposeProduct(
367 const std::vector<std::vector<gate_t>> &contribs,
368 const std::vector<int> &members)
373 std::vector<gate_t> L;
374 for (
int m : members)
for (
gate_t l : contribs[m]) L.push_back(l);
375 std::sort(L.begin(), L.end());
376 L.erase(std::unique(L.begin(), L.end()), L.end());
377 std::map<gate_t, int> idx;
378 for (std::size_t i = 0; i < L.size(); ++i) idx[L[i]] =
static_cast<int>(i);
379 const int nl =
static_cast<int>(L.size());
382 std::vector<std::vector<char>> cooc(nl, std::vector<char>(nl, 0));
383 for (
int m : members) {
384 const auto &cl = contribs[m];
385 for (std::size_t i = 0; i < cl.size(); ++i)
386 for (std::size_t j = i + 1; j < cl.size(); ++j) {
387 int a = idx[cl[i]], b = idx[cl[j]];
388 cooc[a][b] = cooc[b][a] = 1;
394 for (
int u = 0; u < nl; ++u)
395 for (
int v = u + 1; v < nl; ++v)
396 if (!cooc[u][v]) uf.unite(u, v);
397 std::map<int, int> factorId;
398 for (
int u = 0; u < nl; ++u)
399 factorId.emplace(uf.find(u),
static_cast<int>(factorId.size()));
400 const int nf =
static_cast<int>(factorId.size());
401 if (nf < 2)
return out;
403 for (
gate_t l : L) out.leafFactor[l] = factorId[uf.find(idx[l])];
407 std::vector<std::set<std::vector<gate_t>>> parts(nf);
408 for (
int m : members) {
409 std::vector<std::vector<gate_t>> proj(nf);
410 for (
gate_t l : contribs[m])
411 proj[out.leafFactor[l]].push_back(l);
412 for (
int f = 0; f < nf; ++f) {
413 if (proj[f].empty())
return out;
414 parts[f].insert(std::move(proj[f]));
421 std::size_t prod = 1;
422 for (
int f = 0; f < nf; ++f) prod *= parts[f].size();
423 if (prod != members.size())
return out;
425 out.parts.resize(nf);
426 for (
int f = 0; f < nf; ++f)
427 out.parts[f].assign(parts[f].begin(), parts[f].end());
459 std::vector<std::vector<gate_t>> contribs,
462 const std::size_t n = contribs.size();
463 if (n == 0)
return std::vector<double>{1.0};
466 UnionFind uf(
static_cast<int>(n));
468 std::map<gate_t, int> first_owner;
469 for (std::size_t i = 0; i < n; ++i)
470 for (
gate_t l : contribs[i]) {
471 auto it = first_owner.find(l);
472 if (it == first_owner.end()) first_owner[l] =
static_cast<int>(i);
473 else uf.unite(
static_cast<int>(i), it->second);
476 std::map<int, std::vector<int>> blocks;
477 for (std::size_t i = 0; i < n; ++i)
478 blocks[uf.find(
static_cast<int>(i))].push_back(
static_cast<int>(i));
480 std::vector<double> total{1.0};
481 for (
const auto &be : blocks) {
482 const std::vector<int> &members = be.second;
483 std::vector<double> blockPMF;
485 if (members.size() == 1) {
490 blockPMF = std::vector<double>{1.0 - q, q};
492 std::vector<gate_t> common = commonLeaves(contribs, members);
493 if (!common.empty()) {
498 std::vector<double> inner =
499 countPMF(gc, residualsOf(contribs, members, common), ok);
501 blockPMF = std::move(inner);
502 for (
double &x : blockPMF) x *= p_root;
503 blockPMF[0] += (1.0 - p_root);
509 ProductDecomp pd = decomposeProduct(contribs, members);
510 if (!pd.ok) { ok =
false;
return {}; }
511 std::vector<double> acc;
512 for (std::size_t f = 0; f < pd.parts.size(); ++f) {
513 std::vector<double> fp = countPMF(gc, std::move(pd.parts[f]), ok);
515 acc = (f == 0) ? std::move(fp) : productConvolve(acc, fp);
517 blockPMF = std::move(acc);
520 total = convolve(total, blockPMF);
533static double prFromPMF(
const std::vector<double> &pmf,
537 for (
int c = is_scalar ? 0 : 1; c < static_cast<int>(pmf.size()); ++c) {
547 if (sat) pr += pmf[c];
565 std::vector<std::vector<gate_t>> contribs,
bool &ok)
567 if (contribs.empty())
return 1.0;
569 for (
const auto &members : independenceBlocks(contribs)) {
571 if (members.size() == 1) {
574 block_absent = 1.0 - q;
576 std::vector<gate_t> common = commonLeaves(contribs, members);
577 if (!common.empty()) {
580 double inner = pAllAbsent(gc, residualsOf(contribs, members, common), ok);
582 block_absent = (1.0 - p_root) + p_root * inner;
590 ProductDecomp pd = decomposeProduct(contribs, members);
591 if (!pd.ok) { ok =
false;
return 0.0; }
592 double prodPresent = 1.0;
593 for (
auto &fparts : pd.parts) {
594 double fa = pAllAbsent(gc, fparts, ok);
596 prodPresent *= (1.0 - fa);
598 block_absent = 1.0 - prodPresent;
601 result *= block_absent;
613 const std::vector<std::vector<gate_t>> &leaves,
614 const std::vector<long> &vals,
615 const std::vector<std::vector<std::pair<long, double>>> &blocks,
625 auto pAbsentWhere = [&](
auto pred) ->
double {
626 std::vector<std::vector<gate_t>> sub;
627 for (std::size_t i = 0; i < leaves.size(); ++i)
628 if (pred(vals[i])) sub.push_back(leaves[i]);
629 double r = pAllAbsent(gc, std::move(sub), ok);
630 for (
const auto &blk : blocks) {
632 for (
const auto &alt : blk)
if (pred(alt.first)) s += alt.second;
633 r *= 1.0 - (s > 1.0 ? 1.0 : s);
638 const double allAbsent = pAbsentWhere([](
int) {
return true; });
648 - pAbsentWhere([&](
long v){
return v >= C;});
break;
650 - (pAbsentWhere([&](
long v){
return v > C;})
651 - pAbsentWhere([&](
long v){
return v >= C;}));
break;
660 - pAbsentWhere([&](
long v){
return v <= C;});
break;
662 - (pAbsentWhere([&](
long v){
return v < C;})
663 - pAbsentWhere([&](
long v){
return v <= C;}));
break;
675constexpr std::size_t kMaxSumSupport = 1u << 20;
710using JointPMFT = std::map<std::pair<W, long>,
double>;
711using JointPMF = JointPMFT<long>;
721static bool recoverAdditiveSeparation(
722 const std::vector<std::vector<gate_t>> &contribs,
723 const std::vector<int> &members,
const std::vector<long> &weights,
724 const ProductDecomp &pd, std::vector<std::vector<long>> &partVals)
726 const int nf =
static_cast<int>(pd.parts.size());
728 auto partOf = [&](
int m,
int f) {
729 std::vector<gate_t> p;
730 for (
gate_t l : contribs[m])
731 if (pd.leafFactor.at(l) == f) p.push_back(l);
736 std::map<std::vector<std::vector<gate_t>>,
long> grid;
737 for (
int m : members) {
738 std::vector<std::vector<gate_t>> key(nf);
739 for (
int f = 0; f < nf; ++f) key[f] = partOf(m, f);
740 grid[key] = weights[m];
743 const int m0 = members[0];
744 const long W0 = weights[m0];
745 std::vector<std::vector<gate_t>> ref(nf);
746 for (
int f = 0; f < nf; ++f) ref[f] = partOf(m0, f);
749 std::vector<std::map<std::vector<gate_t>,
long>> h(nf);
750 for (
int f = 0; f < nf; ++f)
751 for (
const auto &p : pd.parts[f]) {
752 std::vector<std::vector<gate_t>> key = ref;
754 auto it = grid.find(key);
755 if (it == grid.end())
return false;
756 h[f][p] = it->second - W0;
760 for (
int m : members) {
762 for (
int f = 0; f < nf; ++f) acc += h[f].at(partOf(m, f));
763 if (acc != weights[m])
return false;
766 partVals.assign(nf, {});
767 for (
int f = 0; f < nf; ++f) {
768 partVals[f].reserve(pd.parts[f].size());
769 for (
const auto &p : pd.parts[f])
770 partVals[f].push_back(h[f].at(p) + (f == 0 ? W0 : 0));
777 std::vector<std::vector<gate_t>> contribs,
778 std::vector<W> weights,
bool &ok)
782 if (contribs.empty())
return total;
784 for (
const auto &members : independenceBlocks(contribs)) {
785 JointPMFT<W> blockPMF;
786 if (members.size() == 1) {
789 blockPMF[{0, 0}] += 1.0 - q;
790 blockPMF[{weights[members[0]], 1}] += q;
792 std::vector<gate_t> common = commonLeaves(contribs, members);
793 if (!common.empty()) {
797 std::vector<W> rweights;
798 rweights.reserve(members.size());
799 for (
int m : members) rweights.push_back(weights[m]);
800 JointPMFT<W> inner = sumCountPMF(
801 gc, residualsOf(contribs, members, common), std::move(rweights), ok);
803 for (
const auto &kv : inner) blockPMF[kv.first] += p_root * kv.second;
804 blockPMF[{0, 0}] += 1.0 - p_root;
805 }
else if constexpr (std::is_integral_v<W>) {
813 ProductDecomp pd = decomposeProduct(contribs, members);
814 if (!pd.ok) { ok =
false;
return {}; }
815 std::vector<std::vector<long>> partVals;
816 if (!recoverAdditiveSeparation(contribs, members, weights, pd,
818 ok =
false;
return {};
822 for (std::size_t f = 0; f < pd.parts.size(); ++f) {
823 JointPMFT<W> Jf = sumCountPMF(gc, pd.parts[f], partVals[f], ok);
826 for (
const auto &a : acc)
827 for (
const auto &b : Jf)
828 nacc[{a.first.first * b.first.second
829 + b.first.first * a.first.second,
830 a.first.second * b.first.second}] += a.second * b.second;
831 if (nacc.size() > kMaxSumSupport) { ok =
false;
return {}; }
834 blockPMF = std::move(acc);
838 ok =
false;
return {};
843 for (
const auto &a : total)
844 for (
const auto &b : blockPMF)
845 ntotal[{a.first.first + b.first.first,
846 a.first.second + b.first.second}] += a.second * b.second;
847 if (ntotal.size() > kMaxSumSupport) { ok =
false;
return {}; }
854 std::vector<std::vector<gate_t>> contribs,
855 std::vector<long> weights,
bool &ok);
861static bool i128_mul(__int128 a, __int128 b, __int128 &out)
863 constexpr __int128 LIM =
static_cast<__int128
>(1) << 120;
864 if (a == 0 || b == 0) { out = 0;
return true; }
866 if (r / b != a)
return false;
867 __int128 ar = r < 0 ? -r : r;
868 if (ar > LIM)
return false;
887static std::map<long, double> mulSeparableSumPMF(
889 const std::vector<std::vector<gate_t>> &contribs,
890 const std::vector<int> &members,
const std::vector<long> &weights,
891 const ProductDecomp &pd,
bool &ok)
893 const int nf =
static_cast<int>(pd.parts.size());
895 auto partOf = [&](
int m,
int f) {
896 std::vector<gate_t> p;
897 for (
gate_t l : contribs[m])
898 if (pd.leafFactor.at(l) == f) p.push_back(l);
902 std::map<std::vector<std::vector<gate_t>>,
long> grid;
903 for (
int m : members) {
904 std::vector<std::vector<gate_t>> key(nf);
905 for (
int f = 0; f < nf; ++f) key[f] = partOf(m, f);
906 grid[key] = weights[m];
910 for (
int m : members)
if (weights[m] != 0) { piv = m;
break; }
911 if (piv < 0) { ok =
false;
return {}; }
912 const long D = weights[piv];
913 std::vector<std::vector<gate_t>> ref(nf);
914 for (
int f = 0; f < nf; ++f) ref[f] = partOf(piv, f);
917 std::vector<std::vector<long>> A(nf);
918 std::vector<std::map<std::vector<gate_t>,
int>> partIdx(nf);
919 for (
int f = 0; f < nf; ++f)
920 for (
const auto &p : pd.parts[f]) {
921 partIdx[f][p] =
static_cast<int>(A[f].size());
922 std::vector<std::vector<gate_t>> key = ref;
924 auto it = grid.find(key);
925 if (it == grid.end()) { ok =
false;
return {}; }
926 A[f].push_back(it->second);
931 for (
int t = 0; t < nf - 1; ++t)
932 if (!i128_mul(Dk1,
static_cast<__int128
>(D), Dk1)) { ok =
false;
return {}; }
935 for (
int m : members) {
937 for (
int f = 0; f < nf; ++f) {
938 long a = A[f][partIdx[f].at(partOf(m, f))];
939 if (!i128_mul(prod,
static_cast<__int128
>(a), prod)) { ok =
false;
return {}; }
942 if (!i128_mul(
static_cast<__int128
>(weights[m]), Dk1, rhs)) { ok =
false;
return {}; }
943 if (prod != rhs) { ok =
false;
return {}; }
947 std::map<__int128, double> run;
949 for (
int f = 0; f < nf; ++f) {
950 std::map<long, double> Pf = sumPMF(gc, pd.parts[f], A[f], ok);
952 std::map<__int128, double> nxt;
953 for (
const auto &rk : run)
954 for (
const auto &sk : Pf) {
956 if (!i128_mul(rk.first,
static_cast<__int128
>(sk.first), prod)) {
957 ok =
false;
return {};
959 nxt[prod] += rk.second * sk.second;
961 if (nxt.size() > kMaxSumSupport) { ok =
false;
return {}; }
966 std::map<long, double> out;
967 for (
const auto &kv : run) {
968 if (kv.first % Dk1 != 0) { ok =
false;
return {}; }
969 __int128 bs = kv.first / Dk1;
970 if (bs >
static_cast<__int128
>(LONG_MAX) ||
971 bs <
static_cast<__int128
>(LONG_MIN)) { ok =
false;
return {}; }
972 out[
static_cast<long>(bs)] += kv.second;
985 std::vector<std::vector<gate_t>> contribs,
986 std::vector<long> weights,
bool &ok)
988 std::map<long, double> total;
990 if (contribs.empty())
return total;
992 for (
const auto &members : independenceBlocks(contribs)) {
993 std::map<long, double> blockPMF;
994 if (members.size() == 1) {
997 blockPMF[0] += 1.0 - q;
998 blockPMF[weights[members[0]]] += q;
1000 std::vector<gate_t> common = commonLeaves(contribs, members);
1001 if (!common.empty()) {
1002 double p_root = 1.0;
1005 std::vector<std::vector<gate_t>> residuals =
1006 residualsOf(contribs, members, common);
1007 std::vector<long> rweights;
1008 rweights.reserve(members.size());
1009 for (
int m : members) rweights.push_back(weights[m]);
1011 std::map<long, double> inner =
1012 sumPMF(gc, std::move(residuals), std::move(rweights), ok);
1014 for (
const auto &kv : inner) blockPMF[kv.first] += p_root * kv.second;
1015 blockPMF[0] += (1.0 - p_root);
1027 ProductDecomp pd = decomposeProduct(contribs, members);
1028 if (!pd.ok) { ok =
false;
return {}; }
1029 const int nf =
static_cast<int>(pd.parts.size());
1032 std::map<std::vector<gate_t>,
long> partVal;
1033 for (
int f = 0; f < nf && chosen < 0; ++f) {
1034 std::map<std::vector<gate_t>,
long> pv;
1035 bool consistent =
true;
1036 for (
int m : members) {
1037 std::vector<gate_t> partf;
1038 for (
gate_t l : contribs[m])
1039 if (pd.leafFactor[l] == f) partf.push_back(l);
1040 auto it = pv.find(partf);
1041 if (it == pv.end()) pv[partf] = weights[m];
1042 else if (it->second != weights[m]) { consistent =
false;
break; }
1044 if (consistent) { chosen = f; partVal = std::move(pv); }
1051 std::vector<long> partValues;
1052 partValues.reserve(pd.parts[chosen].size());
1053 for (
const auto &part : pd.parts[chosen])
1054 partValues.push_back(partVal[part]);
1055 std::map<long, double> Sf =
1056 sumPMF(gc, pd.parts[chosen], std::move(partValues), ok);
1060 std::vector<double> M;
1061 for (
int f = 0; f < nf; ++f) {
1062 if (f == chosen)
continue;
1063 std::vector<double> cf = countPMF(gc, pd.parts[f], ok);
1065 M = M.empty() ? std::move(cf) : productConvolve(M, cf);
1069 for (
const auto &skv : Sf)
1070 for (std::size_t mm = 0; mm < M.size(); ++mm)
1072 blockPMF[skv.first *
static_cast<long>(mm)] += skv.second * M[mm];
1073 if (blockPMF.size() > kMaxSumSupport) { ok =
false;
return {}; }
1082 std::vector<std::vector<long>> sep;
1083 if (recoverAdditiveSeparation(contribs, members, weights, pd, sep)) {
1084 std::vector<std::vector<gate_t>> bc;
1085 std::vector<long> bw;
1086 bc.reserve(members.size());
1087 bw.reserve(members.size());
1088 for (
int m : members) {
1089 bc.push_back(contribs[m]);
1090 bw.push_back(weights[m]);
1092 JointPMF j = sumCountPMF(gc, std::move(bc), std::move(bw), ok);
1094 for (
const auto &kv : j) blockPMF[kv.first.first] += kv.second;
1096 blockPMF = mulSeparableSumPMF(gc, contribs, members, weights, pd, ok);
1099 if (blockPMF.size() > kMaxSumSupport) { ok =
false;
return {}; }
1103 std::map<long, double> ntotal;
1104 for (
const auto &a : total)
1105 for (
const auto &b : blockPMF)
1106 ntotal[a.first + b.first] += a.second * b.second;
1107 if (ntotal.size() > kMaxSumSupport) { ok =
false;
return {}; }
1117 unsigned resolved = 0;
1120 std::vector<gate_t> cmps;
1121 for (std::size_t i = 0; i < nb; ++i) {
1122 auto g =
static_cast<gate_t>(i);
1126 if (cmps.empty())
return 0;
1130 for (
gate_t cmp : cmps) {
1145 const auto &ks = match.
ks;
1146 const std::size_t n = ks.size();
1159 std::vector<std::vector<gate_t>> leaves;
1160 std::vector<long> tid_vals;
1163 std::map<gate_t, std::vector<std::pair<double, long>>> blocks;
1164 for (std::size_t i = 0; i < n && ok; ++i) {
1166 const auto &ch = gc.
getWires(ks[i]);
1167 if (ch.size() != 1) { ok =
false;
break; }
1168 blocks[ch[0]].push_back({gc.
getProb(ks[i]),
1169 static_cast<long>(match.
ms[i])});
1171 std::vector<gate_t> ls;
1172 if (parseProductContributor(gc, ks[i], ls)) {
1173 leaves.push_back(std::move(ls));
1174 tid_vals.push_back(
static_cast<long>(match.
ms[i]));
1181 if (!contributorExactMarginal(gc, ks[i], ref, pi)) { ok =
false;
break; }
1182 blocks[ks[i]].push_back({pi,
static_cast<long>(match.
ms[i])});
1193 std::set<gate_t> tidset;
1194 for (
const auto &ls : leaves) tidset.insert(ls.begin(), ls.end());
1196 for (
const auto &b : blocks)
1197 if (tidset.count(b.first)) { clash =
true;
break; }
1198 if (clash)
continue;
1206 if (!aggSubtreePrivate(gc, agg, ref))
continue;
1209 auto blockMass = [](
const std::vector<std::pair<double, long>> &alts) {
1211 for (
const auto &alt : alts) psum += alt.first;
1212 return psum > 1.0 ? 1.0 : psum;
1226 bool all_one =
true;
1227 for (
int m : match.
ms)
if (m != 1) { all_one =
false;
break; }
1228 if (!all_one)
continue;
1229 std::vector<double> total = countPMF(gc, leaves, ok);
1233 for (
const auto &b : blocks) {
1234 double psum = blockMass(b.second);
1235 total = convolve(total, std::vector<double>{1.0 - psum, psum});
1237 const bool is_scalar =
1239 pr = prFromPMF(total, match.
op, match.
C, is_scalar);
1250 auto shift = [&](
long m) {
return is_avg ? m - match.
C : m; };
1251 std::vector<long> weights;
1252 weights.reserve(tid_vals.size());
1253 long lo = 0, hi = 0;
1254 for (
long m : tid_vals) {
1256 weights.push_back(w);
1257 if (w < 0) lo += w;
else hi += w;
1259 for (
const auto &b : blocks)
1260 for (
const auto &alt : b.second) {
1261 long w = shift(alt.second);
1262 if (w < 0) lo += w;
else hi += w;
1265 if (hi - lo + 1 >
static_cast<long>(kMaxSumSupport))
continue;
1266 const long thr = is_avg ? 0 : match.
C;
1268 std::map<long, double> dist = sumPMF(gc, leaves, std::move(weights), ok);
1271 for (
const auto &b : blocks) {
1272 std::map<long, double> bpmf;
1273 for (
const auto &alt : b.second) bpmf[shift(alt.second)] += alt.first;
1274 bpmf[0] += 1.0 - blockMass(b.second);
1275 std::map<long, double> nd;
1276 for (
const auto &a : dist)
1277 for (
const auto &c : bpmf)
1278 nd[a.first + c.first] += a.second * c.second;
1279 if (nd.size() > kMaxSumSupport) { ok =
false;
break; }
1284 for (
const auto &kv : dist)
1285 if (sumSatisfies(kv.first, match.
op, thr)) pr += kv.second;
1290 if (sumSatisfies(0, match.
op, thr)) {
1291 double emptyMass = pAllAbsent(gc, leaves, ok);
1293 for (
const auto &b : blocks) emptyMass *= 1.0 - blockMass(b.second);
1300 std::vector<std::vector<std::pair<long, double>>> blockvec;
1301 blockvec.reserve(blocks.size());
1302 for (
const auto &b : blocks) {
1303 std::vector<std::pair<long, double>> alts;
1304 alts.reserve(b.second.size());
1305 for (
const auto &alt : b.second) alts.push_back({alt.second, alt.first});
1306 blockvec.push_back(std::move(alts));
1308 pr = minMaxProb(gc, leaves, tid_vals, blockvec, agg_kind, match.
op,
1313 if (pr < 0.0) pr = 0.0;
1314 if (pr > 1.0) pr = 1.0;
1331 std::vector<std::vector<gate_t>> contribs;
1332 std::vector<double> values;
1336 if (w.size() != 2)
return 0.0;
1343 std::vector<gate_t> leaves;
1344 if (!parseProductContributor(gc, w[0], leaves))
1346 contribs.push_back(std::move(leaves));
1347 values.push_back(v);
1349 for (
const auto &c : contribs)
1351 if (std::isnan(gc.
getProb(l)))
return 0.0;
1357 JointPMFT<double> pmf = sumCountPMF(gc, std::move(contribs),
1358 std::move(values), pmf_ok);
1359 if (!pmf_ok)
return 0.0;
1364 double num = 0.0, den = 0.0;
1365 for (
const auto &kv : pmf) {
1366 if (kv.first.second < 1)
continue;
1368 num += kv.second * std::pow(kv.first.first
1369 /
static_cast<double>(kv.first.second),
1370 static_cast<double>(k));
1372 if (!(den > 1e-12))
return 0.0;