from fractions import Fraction as F
from pathlib import Path
import itertools,json,math,hashlib,random,re
checks=[]
def eq(name,a,b):
 assert a==b,(name,a,b);checks.append({'name':name,'value':str(a),'expected':str(b),'status':'pass'})
def near(name,a,b,tol=1e-9):
 assert abs(a-b)<=tol,(name,a,b);checks.append({'name':name,'value':a,'expected':b,'tolerance':tol,'status':'pass'})
def mean(x):return sum(x,F(0))/len(x)
def brier(s,p):return p*(1-s)**2+(1-p)*s*s
# Proper score and calibration examples.
for q in [F(0),F(1,4),F(1,2),F(1)]:eq('Brier decomposition p=1/4 q='+str(q),brier(q,F(1,4)),F(3,16)+(q-F(1,4))**2)
eq('coarse ECE cancels',abs(mean([F(2,5),F(1,5)])-mean([F(1,5),F(2,5)])),F(0))
s=[F(1,5),F(4,5)];p=[F(1,10),F(7,10)];pi=mean(p)
rel=mean([(a-b)**2 for a,b in zip(s,p)]);res=mean([(x-pi)**2 for x in p]);unc=pi*(1-pi)
eq('two-bin REL',rel,F(1,100));eq('two-bin RES',res,F(9,100));eq('two-bin UNC',unc,F(6,25));eq('two-bin Brier',mean([brier(a,b) for a,b in zip(s,p)]),F(4,25));eq('decomposition identity',rel-res+unc,F(4,25))
eq('original versus coarsened Brier',mean([F(1,5)**2,(F(2,5)-1)**2]),F(1,5));eq('coarsened Brier',mean([F(3,10)**2,(F(3,10)-1)**2]),F(29,100))
# PAV uses exact rational arithmetic, then checks the certificate for all binary samples up to length 9.
def pav(y,w=None):
 w=w or [F(1)]*len(y);blocks=[]
 for i,(a,b) in enumerate(zip(y,w)):
  blocks.append([i,i+1,b,a*b])
  while len(blocks)>1 and blocks[-2][3]/blocks[-2][2]>blocks[-1][3]/blocks[-1][2]:
   v=blocks.pop();u=blocks.pop();blocks.append([u[0],v[1],u[2]+v[2],u[3]+v[3]])
 q=[F(0)]*len(y)
 for a,b,w,t in blocks:q[a:b]=[t/w]*(b-a)
 return q
cnt=0
for n in range(1,10):
 for yy in itertools.product([F(0),F(1)],repeat=n):
  q=pav(yy);lam=F(0)
  for i in range(n):
   lam+=yy[i]-q[i];assert lam>=0
   if i<n-1:
    assert q[i]<=q[i+1]
    if q[i]<q[i+1]:assert lam==0
  assert lam==0;cnt+=1
checks.append({'name':'PAV cumulative-residual certificate: all binary samples lengths 1..9','cases':cnt,'status':'pass'})
rng=random.Random(8105)
for _ in range(256):
 yy=[F(rng.randrange(11),10) for _ in range(9)];ww=[F(rng.randrange(1,8)) for _ in yy];q=pav(yy,ww);lam=F(0)
 for i in range(len(yy)):
  lam+=ww[i]*(yy[i]-q[i]);assert lam>=0
  if i<len(yy)-1:
   assert q[i]<=q[i+1]
   if q[i]<q[i+1]:assert lam==0
 assert lam==0
checks.append({'name':'PAV weighted rational certificate','cases':256,'status':'pass'})
eq('PAV worked example',pav([F(x) for x in [0,1,0,0,1,1]]),[F(0),F(1,3),F(1,3),F(1,3),F(1),F(1)])
# Temperature and multicalibration.
near('temperature NLL original',-.75*math.log(.9)-.25*math.log(.1),.6546666599918811)
near('temperature NLL T=2',-.75*math.log(.75)-.25*math.log(.25),.5623351446188083)
v=[2/(3+math.sqrt(2)),math.sqrt(2)/(3+math.sqrt(2)),1/(3+math.sqrt(2))]
for a,b in zip(v,[.4530818393219729,.3203772410170408,.22654091966098644]):near('temperature multiclass probability',a,b)
eta=[F(9,10),F(1,2),F(1,10),F(1,2)]
for vals,expected in [([F(1,2)]*4,F(2,25)),([F(7,10),F(7,10),F(1,2),F(1,2)],F(3,50)),([F(7,10),F(2,5),F(1,5),F(1,2)],F(3,200))]:eq('multicalibration potential',mean([(a-b)**2 for a,b in zip(vals,eta)]),expected)
# Reject example and group fairness.
eq('full classification risk',F(1,10)*F(1,20)+F(2,5)*F(3,10)+F(3,10)*F(2,5)+F(1,5)*F(1,20),F(51,200))
eq('accepted-error numerator',F(3,10)*F(1,20),F(3,200));eq('reject total cost',F(3,200)+F(1,5)*F(7,10),F(31,200))
for tp,fn,fp,tn,expected in [(16,4,8,72,(F(4,5),F(1,10),F(6,25),F(2,3))),(48,12,4,36,(F(4,5),F(1,10),F(13,25),F(12,13)))]:eq('fairness confusion table',(F(tp,tp+fn),F(fp,fp+tn),F(tp+fp,100),F(tp,tp+fp)),expected)
# Capstone exact distribution.
eta=[F(1,10),F(2,5),F(4,5),F(1,5),F(3,5),F(9,10)];s=[F(3,20),F(1,2),F(17,20)]*2
cap={}
for g in range(2):cap[f'pi{g}']=mean(eta[3*g:3*g+3])
eq('capstone pi0',cap['pi0'],F(13,30));eq('capstone pi1',cap['pi1'],F(17,30))
eq('capstone Brier S',mean([brier(a,b) for a,b in zip(s,eta)]),F(101,600));eq('capstone Brier full risk',mean([p*(1-p) for p in eta]),F(49,300));eq('capstone improvement',mean([(a-b)**2 for a,b in zip(s,eta)]),F(1,200))
for g,expected in enumerate([(F(12,13),F(8,17),F(3,5)),(F(15,17),F(5,13),F(3,4))]):
 pp=eta[3*g:3*g+3];eq('capstone TPR FPR PPV '+str(g),(sum(pp[1:])/sum(pp),sum(1-p for p in pp[1:])/sum(1-p for p in pp),mean(pp[1:])),expected)
for g,expected in enumerate([(F(81,130),F(49,170)),(F(121,170),F(49,130))]):
 pp=eta[3*g:3*g+3];eq('capstone positive/negative means '+str(g),(sum(p*p for p in pp)/sum(pp),sum(p*(1-p) for p in pp)/sum(1-p for p in pp)),expected)
accept=[i for i,q in enumerate(s) if min(q,1-q)<=F(1,5)];num=sum((eta[i] if s[i]<F(1,2) else 1-eta[i]) for i in accept)/6;cov=F(len(accept),6)
eq('capstone coverage',cov,F(2,3));eq('capstone accepted errors',num,F(1,10));eq('capstone selective risk',num/cov,F(3,20));eq('capstone total reject cost',num+(1-cov)*F(1,5),F(1,6))

print(json.dumps({"status":"pass","numeric_checks":len(checks),"PAV_binary_cases":cnt,"PAV_weighted_cases":256,"results":checks},ensure_ascii=False,indent=2))
