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())
102 const std::vector<unsigned> &ref)
104 std::map<gate_t, unsigned> internalRef;
105 std::set<gate_t> visited;
106 std::vector<gate_t> stk{agg};
108 while (!stk.empty()) {
109 gate_t g = stk.back(); stk.pop_back();
112 if (visited.insert(c).second) stk.push_back(c);
115 for (
gate_t g : visited) {
116 if (g == agg)
continue;
123 if (ref[
static_cast<std::size_t
>(g)] != internalRef[g])
130constexpr unsigned kMaxContributorLeaves = 20;
151 const std::vector<unsigned> &ref,
155 std::map<gate_t, unsigned> internalRef;
156 std::set<gate_t> seen;
157 std::vector<gate_t> order;
158 std::vector<std::pair<gate_t, bool>> stk{{g,
false}};
159 while (!stk.empty()) {
160 auto top = stk.back(); stk.pop_back();
166 if (top.second) { order.push_back(x);
continue; }
167 if (!seen.insert(x).second)
continue;
168 stk.push_back({x,
true});
172 stk.push_back({c,
false});
178 if (x == g)
continue;
183 if (ref[
static_cast<std::size_t
>(x)] != internalRef[x])
return false;
187 const int N =
static_cast<int>(order.size());
188 std::map<gate_t, int> pos;
189 for (
int i = 0; i < N; ++i) pos[order[i]] = i;
190 std::vector<gate_type> typ(N);
191 std::vector<std::vector<int>> childIdx(N);
192 std::vector<int> leafbit(N, -1);
193 std::vector<double> leafProb;
194 for (
int i = 0; i < N; ++i) {
197 leafbit[i] =
static_cast<int>(leafProb.size());
198 leafProb.push_back(gc.
getProb(order[i]));
200 for (
gate_t c : gc.
getWires(order[i])) childIdx[i].push_back(pos[c]);
203 const unsigned m =
static_cast<unsigned>(leafProb.size());
204 if (m > kMaxContributorLeaves)
return false;
207 std::vector<char> val(N);
209 for (uint32_t mask = 0; mask < (1u << m); ++mask) {
210 for (
int i = 0; i < N; ++i) {
214 case gate_input: val[i] = (mask >> leafbit[i]) & 1u;
break;
215 case gate_times: {
char v = 1;
for (
int c : childIdx[i]) v = v && val[c]; val[i] = v;
break; }
216 case gate_plus: {
char v = 0;
for (
int c : childIdx[i]) v = v || val[c]; val[i] = v;
break; }
217 case gate_monus: val[i] = val[childIdx[i][0]] && !val[childIdx[i][1]];
break;
218 default:
return false;
223 for (
unsigned b = 0; b < m; ++b)
224 pr *= (mask >> b) & 1u ? leafProb[b] : 1.0 - leafProb[b];
234 std::vector<int> parent;
235 explicit UnionFind(
int n) : parent(n) {
236 std::iota(parent.begin(), parent.end(), 0);
239 while (parent[x] != x) { parent[x] = parent[parent[x]]; x = parent[x]; }
242 void unite(
int a,
int b) { parent[find(a)] = find(b); }
249static std::vector<std::vector<int>> independenceBlocks(
250 const std::vector<std::vector<gate_t>> &contribs)
252 const std::size_t n = contribs.size();
253 UnionFind uf(
static_cast<int>(n));
254 std::map<gate_t, int> first_owner;
255 for (std::size_t i = 0; i < n; ++i)
256 for (
gate_t l : contribs[i]) {
257 auto it = first_owner.find(l);
258 if (it == first_owner.end()) first_owner[l] =
static_cast<int>(i);
259 else uf.unite(
static_cast<int>(i), it->second);
261 std::map<int, std::vector<int>> bmap;
262 for (std::size_t i = 0; i < n; ++i)
263 bmap[uf.find(
static_cast<int>(i))].push_back(
static_cast<int>(i));
264 std::vector<std::vector<int>> blocks;
265 blocks.reserve(bmap.size());
266 for (
auto &kv : bmap) blocks.push_back(std::move(kv.second));
273static std::vector<gate_t> commonLeaves(
274 const std::vector<std::vector<gate_t>> &contribs,
275 const std::vector<int> &members)
277 std::vector<gate_t> common = contribs[members[0]];
278 for (std::size_t mi = 1; mi < members.size() && !common.empty(); ++mi) {
279 std::vector<gate_t> tmp;
280 std::set_intersection(common.begin(), common.end(),
281 contribs[members[mi]].begin(),
282 contribs[members[mi]].end(),
283 std::back_inserter(tmp));
291static std::vector<std::vector<gate_t>> residualsOf(
292 const std::vector<std::vector<gate_t>> &contribs,
293 const std::vector<int> &members,
const std::vector<gate_t> &common)
295 std::vector<std::vector<gate_t>> residuals;
296 residuals.reserve(members.size());
297 for (
int m : members) {
298 std::vector<gate_t> r;
299 std::set_difference(contribs[m].begin(), contribs[m].end(),
300 common.begin(), common.end(), std::back_inserter(r));
301 residuals.push_back(std::move(r));
307static std::vector<double> convolve(
const std::vector<double> &a,
308 const std::vector<double> &b)
310 if (a.empty())
return b;
311 if (b.empty())
return a;
312 std::vector<double> r(a.size() + b.size() - 1, 0.0);
313 for (std::size_t i = 0; i < a.size(); ++i) {
314 if (a[i] == 0.0)
continue;
315 for (std::size_t j = 0; j < b.size(); ++j)
316 r[i + j] += a[i] * b[j];
324static std::vector<double> productConvolve(
const std::vector<double> &a,
325 const std::vector<double> &b)
327 if (a.empty() || b.empty())
return {};
328 const std::size_t amax = a.size() - 1, bmax = b.size() - 1;
329 std::vector<double> r(amax * bmax + 1, 0.0);
330 for (std::size_t i = 0; i <= amax; ++i) {
331 if (a[i] == 0.0)
continue;
332 for (std::size_t j = 0; j <= bmax; ++j)
333 r[i * j] += a[i] * b[j];
339struct ProductDecomp {
341 std::map<gate_t, int> leafFactor;
342 std::vector<std::vector<std::vector<gate_t>>> parts;
365static ProductDecomp decomposeProduct(
366 const std::vector<std::vector<gate_t>> &contribs,
367 const std::vector<int> &members)
372 std::vector<gate_t> L;
373 for (
int m : members)
for (
gate_t l : contribs[m]) L.push_back(l);
374 std::sort(L.begin(), L.end());
375 L.erase(std::unique(L.begin(), L.end()), L.end());
376 std::map<gate_t, int> idx;
377 for (std::size_t i = 0; i < L.size(); ++i) idx[L[i]] =
static_cast<int>(i);
378 const int nl =
static_cast<int>(L.size());
381 std::vector<std::vector<char>> cooc(nl, std::vector<char>(nl, 0));
382 for (
int m : members) {
383 const auto &cl = contribs[m];
384 for (std::size_t i = 0; i < cl.size(); ++i)
385 for (std::size_t j = i + 1; j < cl.size(); ++j) {
386 int a = idx[cl[i]], b = idx[cl[j]];
387 cooc[a][b] = cooc[b][a] = 1;
393 for (
int u = 0; u < nl; ++u)
394 for (
int v = u + 1; v < nl; ++v)
395 if (!cooc[u][v]) uf.unite(u, v);
396 std::map<int, int> factorId;
397 for (
int u = 0; u < nl; ++u)
398 factorId.emplace(uf.find(u),
static_cast<int>(factorId.size()));
399 const int nf =
static_cast<int>(factorId.size());
400 if (nf < 2)
return out;
402 for (
gate_t l : L) out.leafFactor[l] = factorId[uf.find(idx[l])];
406 std::vector<std::set<std::vector<gate_t>>> parts(nf);
407 for (
int m : members) {
408 std::vector<std::vector<gate_t>> proj(nf);
409 for (
gate_t l : contribs[m])
410 proj[out.leafFactor[l]].push_back(l);
411 for (
int f = 0; f < nf; ++f) {
412 if (proj[f].empty())
return out;
413 parts[f].insert(std::move(proj[f]));
420 std::size_t prod = 1;
421 for (
int f = 0; f < nf; ++f) prod *= parts[f].size();
422 if (prod != members.size())
return out;
424 out.parts.resize(nf);
425 for (
int f = 0; f < nf; ++f)
426 out.parts[f].assign(parts[f].begin(), parts[f].end());
458 std::vector<std::vector<gate_t>> contribs,
461 const std::size_t n = contribs.size();
462 if (n == 0)
return std::vector<double>{1.0};
465 UnionFind uf(
static_cast<int>(n));
467 std::map<gate_t, int> first_owner;
468 for (std::size_t i = 0; i < n; ++i)
469 for (
gate_t l : contribs[i]) {
470 auto it = first_owner.find(l);
471 if (it == first_owner.end()) first_owner[l] =
static_cast<int>(i);
472 else uf.unite(
static_cast<int>(i), it->second);
475 std::map<int, std::vector<int>> blocks;
476 for (std::size_t i = 0; i < n; ++i)
477 blocks[uf.find(
static_cast<int>(i))].push_back(
static_cast<int>(i));
479 std::vector<double> total{1.0};
480 for (
const auto &be : blocks) {
481 const std::vector<int> &members = be.second;
482 std::vector<double> blockPMF;
484 if (members.size() == 1) {
489 blockPMF = std::vector<double>{1.0 - q, q};
491 std::vector<gate_t> common = commonLeaves(contribs, members);
492 if (!common.empty()) {
497 std::vector<double> inner =
498 countPMF(gc, residualsOf(contribs, members, common), ok);
500 blockPMF = std::move(inner);
501 for (
double &x : blockPMF) x *= p_root;
502 blockPMF[0] += (1.0 - p_root);
508 ProductDecomp pd = decomposeProduct(contribs, members);
509 if (!pd.ok) { ok =
false;
return {}; }
510 std::vector<double> acc;
511 for (std::size_t f = 0; f < pd.parts.size(); ++f) {
512 std::vector<double> fp = countPMF(gc, std::move(pd.parts[f]), ok);
514 acc = (f == 0) ? std::move(fp) : productConvolve(acc, fp);
516 blockPMF = std::move(acc);
519 total = convolve(total, blockPMF);
532static double prFromPMF(
const std::vector<double> &pmf,
536 for (
int c = is_scalar ? 0 : 1; c < static_cast<int>(pmf.size()); ++c) {
546 if (sat) pr += pmf[c];
564 std::vector<std::vector<gate_t>> contribs,
bool &ok)
566 if (contribs.empty())
return 1.0;
568 for (
const auto &members : independenceBlocks(contribs)) {
570 if (members.size() == 1) {
573 block_absent = 1.0 - q;
575 std::vector<gate_t> common = commonLeaves(contribs, members);
576 if (!common.empty()) {
579 double inner = pAllAbsent(gc, residualsOf(contribs, members, common), ok);
581 block_absent = (1.0 - p_root) + p_root * inner;
589 ProductDecomp pd = decomposeProduct(contribs, members);
590 if (!pd.ok) { ok =
false;
return 0.0; }
591 double prodPresent = 1.0;
592 for (
auto &fparts : pd.parts) {
593 double fa = pAllAbsent(gc, fparts, ok);
595 prodPresent *= (1.0 - fa);
597 block_absent = 1.0 - prodPresent;
600 result *= block_absent;
612 const std::vector<std::vector<gate_t>> &leaves,
613 const std::vector<long> &vals,
614 const std::vector<std::vector<std::pair<long, double>>> &blocks,
624 auto pAbsentWhere = [&](
auto pred) ->
double {
625 std::vector<std::vector<gate_t>> sub;
626 for (std::size_t i = 0; i < leaves.size(); ++i)
627 if (pred(vals[i])) sub.push_back(leaves[i]);
628 double r = pAllAbsent(gc, std::move(sub), ok);
629 for (
const auto &blk : blocks) {
631 for (
const auto &alt : blk)
if (pred(alt.first)) s += alt.second;
632 r *= 1.0 - (s > 1.0 ? 1.0 : s);
637 const double allAbsent = pAbsentWhere([](
int) {
return true; });
647 - pAbsentWhere([&](
long v){
return v >= C;});
break;
649 - (pAbsentWhere([&](
long v){
return v > C;})
650 - pAbsentWhere([&](
long v){
return v >= C;}));
break;
659 - pAbsentWhere([&](
long v){
return v <= C;});
break;
661 - (pAbsentWhere([&](
long v){
return v < C;})
662 - pAbsentWhere([&](
long v){
return v <= C;}));
break;
674constexpr std::size_t kMaxSumSupport = 1u << 20;
709using JointPMFT = std::map<std::pair<W, long>,
double>;
710using JointPMF = JointPMFT<long>;
720static bool recoverAdditiveSeparation(
721 const std::vector<std::vector<gate_t>> &contribs,
722 const std::vector<int> &members,
const std::vector<long> &weights,
723 const ProductDecomp &pd, std::vector<std::vector<long>> &partVals)
725 const int nf =
static_cast<int>(pd.parts.size());
727 auto partOf = [&](
int m,
int f) {
728 std::vector<gate_t> p;
729 for (
gate_t l : contribs[m])
730 if (pd.leafFactor.at(l) == f) p.push_back(l);
735 std::map<std::vector<std::vector<gate_t>>,
long> grid;
736 for (
int m : members) {
737 std::vector<std::vector<gate_t>> key(nf);
738 for (
int f = 0; f < nf; ++f) key[f] = partOf(m, f);
739 grid[key] = weights[m];
742 const int m0 = members[0];
743 const long W0 = weights[m0];
744 std::vector<std::vector<gate_t>> ref(nf);
745 for (
int f = 0; f < nf; ++f) ref[f] = partOf(m0, f);
748 std::vector<std::map<std::vector<gate_t>,
long>> h(nf);
749 for (
int f = 0; f < nf; ++f)
750 for (
const auto &p : pd.parts[f]) {
751 std::vector<std::vector<gate_t>> key = ref;
753 auto it = grid.find(key);
754 if (it == grid.end())
return false;
755 h[f][p] = it->second - W0;
759 for (
int m : members) {
761 for (
int f = 0; f < nf; ++f) acc += h[f].at(partOf(m, f));
762 if (acc != weights[m])
return false;
765 partVals.assign(nf, {});
766 for (
int f = 0; f < nf; ++f) {
767 partVals[f].reserve(pd.parts[f].size());
768 for (
const auto &p : pd.parts[f])
769 partVals[f].push_back(h[f].at(p) + (f == 0 ? W0 : 0));
776 std::vector<std::vector<gate_t>> contribs,
777 std::vector<W> weights,
bool &ok)
781 if (contribs.empty())
return total;
783 for (
const auto &members : independenceBlocks(contribs)) {
784 JointPMFT<W> blockPMF;
785 if (members.size() == 1) {
788 blockPMF[{0, 0}] += 1.0 - q;
789 blockPMF[{weights[members[0]], 1}] += q;
791 std::vector<gate_t> common = commonLeaves(contribs, members);
792 if (!common.empty()) {
796 std::vector<W> rweights;
797 rweights.reserve(members.size());
798 for (
int m : members) rweights.push_back(weights[m]);
799 JointPMFT<W> inner = sumCountPMF(
800 gc, residualsOf(contribs, members, common), std::move(rweights), ok);
802 for (
const auto &kv : inner) blockPMF[kv.first] += p_root * kv.second;
803 blockPMF[{0, 0}] += 1.0 - p_root;
804 }
else if constexpr (std::is_integral_v<W>) {
812 ProductDecomp pd = decomposeProduct(contribs, members);
813 if (!pd.ok) { ok =
false;
return {}; }
814 std::vector<std::vector<long>> partVals;
815 if (!recoverAdditiveSeparation(contribs, members, weights, pd,
817 ok =
false;
return {};
821 for (std::size_t f = 0; f < pd.parts.size(); ++f) {
822 JointPMFT<W> Jf = sumCountPMF(gc, pd.parts[f], partVals[f], ok);
825 for (
const auto &a : acc)
826 for (
const auto &b : Jf)
827 nacc[{a.first.first * b.first.second
828 + b.first.first * a.first.second,
829 a.first.second * b.first.second}] += a.second * b.second;
830 if (nacc.size() > kMaxSumSupport) { ok =
false;
return {}; }
833 blockPMF = std::move(acc);
837 ok =
false;
return {};
842 for (
const auto &a : total)
843 for (
const auto &b : blockPMF)
844 ntotal[{a.first.first + b.first.first,
845 a.first.second + b.first.second}] += a.second * b.second;
846 if (ntotal.size() > kMaxSumSupport) { ok =
false;
return {}; }
853 std::vector<std::vector<gate_t>> contribs,
854 std::vector<long> weights,
bool &ok);
860static bool i128_mul(__int128 a, __int128 b, __int128 &out)
862 constexpr __int128 LIM =
static_cast<__int128
>(1) << 120;
863 if (a == 0 || b == 0) { out = 0;
return true; }
865 if (r / b != a)
return false;
866 __int128 ar = r < 0 ? -r : r;
867 if (ar > LIM)
return false;
886static std::map<long, double> mulSeparableSumPMF(
888 const std::vector<std::vector<gate_t>> &contribs,
889 const std::vector<int> &members,
const std::vector<long> &weights,
890 const ProductDecomp &pd,
bool &ok)
892 const int nf =
static_cast<int>(pd.parts.size());
894 auto partOf = [&](
int m,
int f) {
895 std::vector<gate_t> p;
896 for (
gate_t l : contribs[m])
897 if (pd.leafFactor.at(l) == f) p.push_back(l);
901 std::map<std::vector<std::vector<gate_t>>,
long> grid;
902 for (
int m : members) {
903 std::vector<std::vector<gate_t>> key(nf);
904 for (
int f = 0; f < nf; ++f) key[f] = partOf(m, f);
905 grid[key] = weights[m];
909 for (
int m : members)
if (weights[m] != 0) { piv = m;
break; }
910 if (piv < 0) { ok =
false;
return {}; }
911 const long D = weights[piv];
912 std::vector<std::vector<gate_t>> ref(nf);
913 for (
int f = 0; f < nf; ++f) ref[f] = partOf(piv, f);
916 std::vector<std::vector<long>> A(nf);
917 std::vector<std::map<std::vector<gate_t>,
int>> partIdx(nf);
918 for (
int f = 0; f < nf; ++f)
919 for (
const auto &p : pd.parts[f]) {
920 partIdx[f][p] =
static_cast<int>(A[f].size());
921 std::vector<std::vector<gate_t>> key = ref;
923 auto it = grid.find(key);
924 if (it == grid.end()) { ok =
false;
return {}; }
925 A[f].push_back(it->second);
930 for (
int t = 0; t < nf - 1; ++t)
931 if (!i128_mul(Dk1,
static_cast<__int128
>(D), Dk1)) { ok =
false;
return {}; }
934 for (
int m : members) {
936 for (
int f = 0; f < nf; ++f) {
937 long a = A[f][partIdx[f].at(partOf(m, f))];
938 if (!i128_mul(prod,
static_cast<__int128
>(a), prod)) { ok =
false;
return {}; }
941 if (!i128_mul(
static_cast<__int128
>(weights[m]), Dk1, rhs)) { ok =
false;
return {}; }
942 if (prod != rhs) { ok =
false;
return {}; }
946 std::map<__int128, double> run;
948 for (
int f = 0; f < nf; ++f) {
949 std::map<long, double> Pf = sumPMF(gc, pd.parts[f], A[f], ok);
951 std::map<__int128, double> nxt;
952 for (
const auto &rk : run)
953 for (
const auto &sk : Pf) {
955 if (!i128_mul(rk.first,
static_cast<__int128
>(sk.first), prod)) {
956 ok =
false;
return {};
958 nxt[prod] += rk.second * sk.second;
960 if (nxt.size() > kMaxSumSupport) { ok =
false;
return {}; }
965 std::map<long, double> out;
966 for (
const auto &kv : run) {
967 if (kv.first % Dk1 != 0) { ok =
false;
return {}; }
968 __int128 bs = kv.first / Dk1;
969 if (bs >
static_cast<__int128
>(LONG_MAX) ||
970 bs <
static_cast<__int128
>(LONG_MIN)) { ok =
false;
return {}; }
971 out[
static_cast<long>(bs)] += kv.second;
984 std::vector<std::vector<gate_t>> contribs,
985 std::vector<long> weights,
bool &ok)
987 std::map<long, double> total;
989 if (contribs.empty())
return total;
991 for (
const auto &members : independenceBlocks(contribs)) {
992 std::map<long, double> blockPMF;
993 if (members.size() == 1) {
996 blockPMF[0] += 1.0 - q;
997 blockPMF[weights[members[0]]] += q;
999 std::vector<gate_t> common = commonLeaves(contribs, members);
1000 if (!common.empty()) {
1001 double p_root = 1.0;
1004 std::vector<std::vector<gate_t>> residuals =
1005 residualsOf(contribs, members, common);
1006 std::vector<long> rweights;
1007 rweights.reserve(members.size());
1008 for (
int m : members) rweights.push_back(weights[m]);
1010 std::map<long, double> inner =
1011 sumPMF(gc, std::move(residuals), std::move(rweights), ok);
1013 for (
const auto &kv : inner) blockPMF[kv.first] += p_root * kv.second;
1014 blockPMF[0] += (1.0 - p_root);
1026 ProductDecomp pd = decomposeProduct(contribs, members);
1027 if (!pd.ok) { ok =
false;
return {}; }
1028 const int nf =
static_cast<int>(pd.parts.size());
1031 std::map<std::vector<gate_t>,
long> partVal;
1032 for (
int f = 0; f < nf && chosen < 0; ++f) {
1033 std::map<std::vector<gate_t>,
long> pv;
1034 bool consistent =
true;
1035 for (
int m : members) {
1036 std::vector<gate_t> partf;
1037 for (
gate_t l : contribs[m])
1038 if (pd.leafFactor[l] == f) partf.push_back(l);
1039 auto it = pv.find(partf);
1040 if (it == pv.end()) pv[partf] = weights[m];
1041 else if (it->second != weights[m]) { consistent =
false;
break; }
1043 if (consistent) { chosen = f; partVal = std::move(pv); }
1050 std::vector<long> partValues;
1051 partValues.reserve(pd.parts[chosen].size());
1052 for (
const auto &part : pd.parts[chosen])
1053 partValues.push_back(partVal[part]);
1054 std::map<long, double> Sf =
1055 sumPMF(gc, pd.parts[chosen], std::move(partValues), ok);
1059 std::vector<double> M;
1060 for (
int f = 0; f < nf; ++f) {
1061 if (f == chosen)
continue;
1062 std::vector<double> cf = countPMF(gc, pd.parts[f], ok);
1064 M = M.empty() ? std::move(cf) : productConvolve(M, cf);
1068 for (
const auto &skv : Sf)
1069 for (std::size_t mm = 0; mm < M.size(); ++mm)
1071 blockPMF[skv.first *
static_cast<long>(mm)] += skv.second * M[mm];
1072 if (blockPMF.size() > kMaxSumSupport) { ok =
false;
return {}; }
1081 std::vector<std::vector<long>> sep;
1082 if (recoverAdditiveSeparation(contribs, members, weights, pd, sep)) {
1083 std::vector<std::vector<gate_t>> bc;
1084 std::vector<long> bw;
1085 bc.reserve(members.size());
1086 bw.reserve(members.size());
1087 for (
int m : members) {
1088 bc.push_back(contribs[m]);
1089 bw.push_back(weights[m]);
1091 JointPMF j = sumCountPMF(gc, std::move(bc), std::move(bw), ok);
1093 for (
const auto &kv : j) blockPMF[kv.first.first] += kv.second;
1095 blockPMF = mulSeparableSumPMF(gc, contribs, members, weights, pd, ok);
1098 if (blockPMF.size() > kMaxSumSupport) { ok =
false;
return {}; }
1102 std::map<long, double> ntotal;
1103 for (
const auto &a : total)
1104 for (
const auto &b : blockPMF)
1105 ntotal[a.first + b.first] += a.second * b.second;
1106 if (ntotal.size() > kMaxSumSupport) { ok =
false;
return {}; }
1116 unsigned resolved = 0;
1119 std::vector<gate_t> cmps;
1120 for (std::size_t i = 0; i < nb; ++i) {
1121 auto g =
static_cast<gate_t>(i);
1125 if (cmps.empty())
return 0;
1129 for (
gate_t cmp : cmps) {
1144 const auto &ks = match.
ks;
1145 const std::size_t n = ks.size();
1149 if (ref[
static_cast<std::size_t
>(agg)] != 1)
continue;
1158 std::vector<std::vector<gate_t>> leaves;
1159 std::vector<long> tid_vals;
1162 std::map<gate_t, std::vector<std::pair<double, long>>> blocks;
1163 for (std::size_t i = 0; i < n && ok; ++i) {
1165 const auto &ch = gc.
getWires(ks[i]);
1166 if (ch.size() != 1) { ok =
false;
break; }
1167 blocks[ch[0]].push_back({gc.
getProb(ks[i]),
1168 static_cast<long>(match.
ms[i])});
1170 std::vector<gate_t> ls;
1171 if (parseProductContributor(gc, ks[i], ls)) {
1172 leaves.push_back(std::move(ls));
1173 tid_vals.push_back(
static_cast<long>(match.
ms[i]));
1180 if (!contributorExactMarginal(gc, ks[i], ref, pi)) { ok =
false;
break; }
1181 blocks[ks[i]].push_back({pi,
static_cast<long>(match.
ms[i])});
1192 std::set<gate_t> tidset;
1193 for (
const auto &ls : leaves) tidset.insert(ls.begin(), ls.end());
1195 for (
const auto &b : blocks)
1196 if (tidset.count(b.first)) { clash =
true;
break; }
1197 if (clash)
continue;
1205 if (!aggSubtreePrivate(gc, agg, ref))
continue;
1208 auto blockMass = [](
const std::vector<std::pair<double, long>> &alts) {
1210 for (
const auto &alt : alts) psum += alt.first;
1211 return psum > 1.0 ? 1.0 : psum;
1225 bool all_one =
true;
1226 for (
int m : match.
ms)
if (m != 1) { all_one =
false;
break; }
1227 if (!all_one)
continue;
1228 std::vector<double> total = countPMF(gc, leaves, ok);
1232 for (
const auto &b : blocks) {
1233 double psum = blockMass(b.second);
1234 total = convolve(total, std::vector<double>{1.0 - psum, psum});
1236 const bool is_scalar =
1238 pr = prFromPMF(total, match.
op, match.
C, is_scalar);
1249 auto shift = [&](
long m) {
return is_avg ? m - match.
C : m; };
1250 std::vector<long> weights;
1251 weights.reserve(tid_vals.size());
1252 long lo = 0, hi = 0;
1253 for (
long m : tid_vals) {
1255 weights.push_back(w);
1256 if (w < 0) lo += w;
else hi += w;
1258 for (
const auto &b : blocks)
1259 for (
const auto &alt : b.second) {
1260 long w = shift(alt.second);
1261 if (w < 0) lo += w;
else hi += w;
1264 if (hi - lo + 1 >
static_cast<long>(kMaxSumSupport))
continue;
1265 const long thr = is_avg ? 0 : match.
C;
1267 std::map<long, double> dist = sumPMF(gc, leaves, std::move(weights), ok);
1270 for (
const auto &b : blocks) {
1271 std::map<long, double> bpmf;
1272 for (
const auto &alt : b.second) bpmf[shift(alt.second)] += alt.first;
1273 bpmf[0] += 1.0 - blockMass(b.second);
1274 std::map<long, double> nd;
1275 for (
const auto &a : dist)
1276 for (
const auto &c : bpmf)
1277 nd[a.first + c.first] += a.second * c.second;
1278 if (nd.size() > kMaxSumSupport) { ok =
false;
break; }
1283 for (
const auto &kv : dist)
1284 if (sumSatisfies(kv.first, match.
op, thr)) pr += kv.second;
1289 if (sumSatisfies(0, match.
op, thr)) {
1290 double emptyMass = pAllAbsent(gc, leaves, ok);
1292 for (
const auto &b : blocks) emptyMass *= 1.0 - blockMass(b.second);
1299 std::vector<std::vector<std::pair<long, double>>> blockvec;
1300 blockvec.reserve(blocks.size());
1301 for (
const auto &b : blocks) {
1302 std::vector<std::pair<long, double>> alts;
1303 alts.reserve(b.second.size());
1304 for (
const auto &alt : b.second) alts.push_back({alt.second, alt.first});
1305 blockvec.push_back(std::move(alts));
1307 pr = minMaxProb(gc, leaves, tid_vals, blockvec, agg_kind, match.
op,
1312 if (pr < 0.0) pr = 0.0;
1313 if (pr > 1.0) pr = 1.0;
1330 std::vector<std::vector<gate_t>> contribs;
1331 std::vector<double> values;
1335 if (w.size() != 2)
return 0.0;
1342 std::vector<gate_t> leaves;
1343 if (!parseProductContributor(gc, w[0], leaves))
1345 contribs.push_back(std::move(leaves));
1346 values.push_back(v);
1348 for (
const auto &c : contribs)
1350 if (std::isnan(gc.
getProb(l)))
return 0.0;
1356 JointPMFT<double> pmf = sumCountPMF(gc, std::move(contribs),
1357 std::move(values), pmf_ok);
1358 if (!pmf_ok)
return 0.0;
1363 double num = 0.0, den = 0.0;
1364 for (
const auto &kv : pmf) {
1365 if (kv.first.second < 1)
continue;
1367 num += kv.second * std::pow(kv.first.first
1368 /
static_cast<double>(kv.first.second),
1369 static_cast<double>(k));
1371 if (!(den > 1e-12))
return 0.0;