#!/usr/bin/env python3
"""L3b可复算预算核验。只用Python标准库。
用法: python foundations-simulation-budget-check.py [--output results.json] [--simulate]
有理数/有限枚举是精确检查；重复模拟仅诊断实现，不构成误差定理。
不会下载、联网、修改原始项目或调用外部程序。
"""
from fractions import Fraction as F
from itertools import product, combinations
from math import comb, prod, sqrt, sin, pi, ceil
from functools import lru_cache
from statistics import mean, variance
from pathlib import Path
import argparse,json,random
checks=[];outputs={}
def check(label,a,b):
 if a!=b:raise ValueError((label,a,b))
 checks.append({'check':label,'value':str(a),'status':'passed'})
def record(k,v):outputs[k]=v
# RQMC whole-net arithmetic: do not treat grid rows as independent.
Q=[sum(((F(i,8)+d)%1) for i in range(8))/8 for d in [F(1,64),F(11,64),F(21,64),F(31,64)]]
check('rqmc_Q_values',Q,[F(29,64),F(31,64),F(33,64),F(35,64)])
check('rqmc_pooled_mean',mean(Q),F(1,2));check('rqmc_estimated_variance',variance(Q)/4,F(5,12288))
for n,R,vs,vj in [(8,4,F(1,3072),F(1,24576)),(8,16,F(1,12288),F(1,98304)),(16,8,F(1,24576),F(1,393216))]:
 check(f'shift_variance_{n}_{R}',F(1,12*R*n*n),vs);check(f'jitter_variance_{n}_{R}',F(1,12*R*n**3),vj)
# Fourier cancellation is exact for finite trigonometric polynomials; numerical
# evaluation here only confirms the implementation's trigonometric path.
max_alias_error=max(abs(mean(sin(2*pi*8*((i/8+d)%1)) for i in range(8))-sin(2*pi*8*d)) for d in [.013,.141,.283,.777])
if not max_alias_error<1e-12:raise ValueError(("Fourier alias roundoff",max_alias_error))
record('rqmc_alias_roundoff',max_alias_error)
# MLMC oracle allocation and fully rational realized variance.
V=[F(1,4**l) for l in range(5)];C=[2**l for l in range(5)];S=sum(sqrt(float(v*c)) for v,c in zip(V,C))
N=[ceil(256*S*sqrt(float(v/c))) for v,c in zip(V,C)]
check('mlmc_integer_allocation',N,[720,255,90,32,12])
check('mlmc_cost',sum(n*c for n,c in zip(N,C)),2038)
mlv=sum(v/n for v,n in zip(V,N));check('mlmc_variance',mlv,F(135,34816));check('mlmc_mse',mlv+F(1,256),F(271,34816))
check('mlmc_direct_variance',sum(V),F(341,256));check('mlmc_direct_budget_cost',341*16,5456)
check('capstone_mlmc_variance',F(1,112)+F(1,224)+F(1,512),F(55,3584));check('capstone_mlmc_mse',F(55,3584)+F(1,256),F(69,3584))
record('mlmc',{'N':N,'variance':str(mlv),'mse':str(mlv+F(1,256)),'cost':2038})
# Gaussian moments, evaluated algebraically rather than sampling.
def normalmoment(k):return 0 if k%2 else prod(range(1,k,2))
def shiftednormalmoment(k,mu,var):return sum(F(comb(k,j))*mu**(k-j)*var**(j//2)*normalmoment(j) for j in range(0,k+1,2))
for M,N0,expected in [(8,512,F(337,16384)),(16,256,F(417,32768)),(32,128,F(1153,65536))]:
 a=1+F(1,M);ew2=shiftednormalmoment(2,F(0),a);ew4=shiftednormalmoment(4,F(0),a)
 check(f'nested_square_mse_M{M}',(ew4-ew2**2)/N0+(ew2-1)**2,expected)
record('nested_square_continuous_optimum_condition','M^3-M=T; derivative computed in prose')
for M in [1,2,4,8,16,32]:
 ez3=sum(shiftednormalmoment(3,F(y),F(1,M)) for y in [0,1])/2
 ez6=sum(shiftednormalmoment(6,F(y),F(1,M)) for y in [0,1])/2
 check(f'cubic_bias_M{M}',ez3-F(1,2),F(3,2*M))
 check(f'cubic_variance_M{M}',ez6-ez3**2,F(1,4)+F(6,M)+F(81,4*M*M)+F(15,M**3))
 if M==16:record('capstone_cubic_plugin_mse',str((ez6-ez3**2)/192+F(9,4*M*M)))
check('square_independent_product_variance',F(3)+2+1-1,5)
check('cubic_independent_product_variance',F(1+8,2)-F(1,4),F(17,4))
check('H3_fourth_moment',sum(comb(4,j)*(-3)**(4-j)*normalmoment(2*j+4) for j in range(5)),3348)
# Generic finite product-space Hoeffding construction, including biased input.
def test_projections(m,n,p,h,label):
 vals=[-1,1];pr={-1:1-p,1:p}
 def prob(xs):return prod(pr[x] for x in xs)
 @lru_cache(None)
 def g(xs):return sum(prob(t)*h(xs+t) for t in product(vals,repeat=m-len(xs)))
 @lru_cache(None)
 def hs(xs):
  s=len(xs)
  return sum((-1)**(s-len(A))*g(tuple(xs[i] for i in A)) for k in range(s+1) for A in combinations(range(s),k))
 theta=g(());sig=[]
 for s in range(1,m+1):
  sig.append(sum(prob(xs)*hs(xs)**2 for xs in product(vals,repeat=s)))
  for xs in product(vals,repeat=s-1):check(f'{label}:canonical_s{s}_{xs}',sum(pr[z]*hs(xs+(z,)) for z in vals),0)
 for xs in product(vals,repeat=m):
  total=sum(hs(tuple(xs[i] for i in A)) for s in range(m+1) for A in combinations(range(m),s))
  check(f'{label}:reconstruct_{xs}',total,h(xs))
 predicted=sum(F(comb(m,s)**2,comb(n,s))*sig[s-1] for s in range(1,m+1))
 e=F(0);e2=F(0);kernel_e2=sum(prob(xs)*h(xs)**2 for xs in product(vals,repeat=m))
 K=comb(n,m)
 for xs in product(vals,repeat=n):
  table=[h(tuple(xs[i] for i in I)) for I in combinations(range(n),m)]
  u=sum(table)/K;probx=prob(xs);e+=probx*u;e2+=probx*u*u
  tau=sum((z-u)**2 for z in table)/K
  # Enumerate all ordered B=2 draws; conditional samples may repeat.
  draws=[(a+b)/2 for a,b in product(table,repeat=2)]
  check(f'{label}:conditional_wr_{xs}',sum((z-u)**2 for z in draws)/(K*K),tau/2)
  draws_nr=[(a+b)/2 for a,b in combinations(table,2)]
  check(f'{label}:conditional_wor_{xs}',sum((z-u)**2 for z in draws_nr)/comb(K,2),F(K-2,2*(K-1))*tau)
 check(f'{label}:mean',e,theta);check(f'{label}:variance',e2-e*e,predicted)
 check(f'{label}:kernel_variance',kernel_e2-theta**2,sum(comb(m,s)*sig[s-1] for s in range(1,m+1)))
 return theta,predicted,kernel_e2-theta**2,sig
for m in [2,3,4]:
 for p in [F(1,2),F(3,4)]:test_projections(m,5,p,lambda xs:F(prod(xs)),f'product_m{m}_p{p}')
theta,v,sv,sig=test_projections(3,6,F(3,4),lambda xs:F(prod(xs)),'capstone_biased_rank')
check('capstone_projection_variances',sig,[F(3,64),F(9,64),F(27,64)])
check('capstone_full_U_variance',v,F(45,256));check('capstone_incomplete_wr',v+(sv-v)/10,F(657,2560));check('capstone_incomplete_wor',v+(sv-v)/19,F(531,2432))
test_projections(3,6,F(1,2),lambda xs:F(sum(xs)+2*sum(a*b for a,b in combinations(xs,2))+3*prod(xs)),'mixed_123')
for n in range(3,9):
 e2=F(0)
 for xs in product([-1,1],repeat=n):
  u=F(sum(prod(xs[i] for i in I) for I in combinations(range(n),3)),comb(n,3));S=sum(xs)
  check(f'H3_identity_n{n}_{xs}',u,F(S**3-(3*n-2)*S,n*(n-1)*(n-2)));e2+=u*u/F(2**n)
 check(f'H3_variance_n{n}',e2,F(1,comb(n,3)))
check('incomplete_n100_linear_B100',F(99,100)*F(1,100)+F(1,200),F(149,10000))
check('incomplete_n100_degenerate_B100',F(99,100)*F(1,4950)+F(1,100),F(51,5000))
check('incomplete_rank3_wr',F(5,6)*F(1,20)+F(1,6),F(5,24))
check('incomplete_rank3_wor',F(1,20)+F(20-6,6*19)*(1-F(1,20)),F(1,6))
# Same deterministic seed makes diagnostics reproducible, not certain.
def simulate():
 rng=random.Random(20261008);R=4;n=8;B=10000
 q=[mean(mean(((i/n+rngdelta)%1) for i in range(n)) for rngdelta in [rng.random() for _ in range(R)]) for _ in range(B)]
 est=variance(q);truth=1/(12*R*n*n)
 out={'rqmc':{'runs':B,'empirical_variance':est,'theory':truth,'ratio':est/truth}}
 N0=16;M=8;runs=4000
 samples=[]
 for _ in range(runs):
  ws=[]
  for i in range(N0):
   y=rng.gauss(0,1);z=mean(y+rng.gauss(0,1) for j in range(M));ws.append(z*z)
  samples.append(mean(ws))
 truthvar=2*(1+1/M)**2/N0
 out['nested_square']={'runs':runs,'empirical_mean':mean(samples),'theory_mean':1+1/M,'empirical_variance':variance(samples),'theory_variance':truthvar,'variance_ratio':variance(samples)/truthvar}
 out['interpretation']='Finite-run diagnostics only; exact identities above are the mathematical checks.'
 return out
if __name__=='__main__':
 parser=argparse.ArgumentParser(description=__doc__);parser.add_argument('--output');parser.add_argument('--simulate',action='store_true');args=parser.parse_args()
 result={'status':'passed','exact_check_count':len(checks),'outputs':outputs,'checks':checks}
 if args.simulate:result['simulation_diagnostics']=simulate()
 data=json.dumps(result,ensure_ascii=False,indent=2)
 if args.output:Path(args.output).write_text(data+'\n')
 print(json.dumps({k:v for k,v in result.items() if k!='checks'},ensure_ascii=False,indent=2))
