#!/usr/bin/env python3
"""Recompute binary-likelihood examples. Python 3 + mpmath; no repository imports.

Exact finite-distribution checks use Fraction. Display decimals use 90 digits.
--output writes the JSON certificate; omitted output prints it to stdout.
These finite checks supplement, rather than replace, the proofs in the pages.
"""
import argparse
import itertools as it
import json
from fractions import Fraction as F
from math import comb, floor
from pathlib import Path
import mpmath as mp
mp.mp.dps = 90
COUNTS = {}

def check(ok, group):
    if not ok:
        raise AssertionError(group)
    COUNTS[group] = COUNTS.get(group, 0) + 1

def close(a, b, group, tol=mp.mpf('1e-75')):
    check(abs(a-b) <= tol * (1+abs(a)+abs(b)), group)

def conv(a, b):
    c = [0]*(len(a)+len(b)-1)
    for i, x in enumerate(a):
        for j, y in enumerate(b):
            c[i+j] += x*y
    return c

def det(a):
    a = [list(map(F, r)) for r in a]; n=len(a); s=F(1)
    for k in range(n):
        pivot = next((i for i in range(k,n) if a[i][k]), None)
        if pivot is None: return F(0)
        if pivot != k: a[k], a[pivot] = a[pivot], a[k]; s=-s
        d=a[k][k]; s*=d
        for i in range(k+1,n):
            q=a[i][k]/d
            for j in range(k+1,n): a[i][j]-=q*a[k][j]
    return s

def matinfo(x, w):
    d=len(x[0])
    return [[sum(F(w[k])*x[k][i]*x[k][j] for k in range(len(x)))
             for j in range(d)] for i in range(d)]

def dp_moments(xs, ws, m):
    d=len(xs[0]); z=[F(0)]*(m+1); z[0]=F(1)
    a=[[F(0)]*d for _ in range(m+1)]
    b=[[[F(0)]*d for _ in range(d)] for _ in range(m+1)]
    for k,(x,w) in enumerate(zip(xs,ws),1):
        for r in range(min(k,m),0,-1):
            zm=z[r-1]; am=a[r-1]; bm=b[r-1]
            b[r]=[[b[r][i][j]+w*(bm[i][j]+x[i]*am[j]+am[i]*x[j]+x[i]*x[j]*zm)
                   for j in range(d)] for i in range(d)]
            a[r]=[a[r][i]+w*(am[i]+x[i]*zm) for i in range(d)]
            z[r]+=w*zm
    mean=[v/z[m] for v in a[m]]
    cov=[[b[m][i][j]/z[m]-mean[i]*mean[j] for j in range(d)] for i in range(d)]
    return z[m],mean,cov

def enumerate_moments(xs,ws,m):
    d=len(xs[0]); rows=[]
    for s in it.combinations(range(len(xs)),m):
        w=F(1)
        for i in s: w*=ws[i]
        t=[sum(xs[i][j] for i in s) for j in range(d)]
        rows.append((s,t,w))
    z=sum(r[2] for r in rows)
    mean=[sum(w*t[i] for s,t,w in rows)/z for i in range(d)]
    cov=[[sum(w*t[i]*t[j] for s,t,w in rows)/z-mean[i]*mean[j]
          for j in range(d)] for i in range(d)]
    return z,mean,cov,rows

def table_coeff(n1,n0,k):
    lo=max(0,k-n0); hi=min(n1,k)
    return lo,[comb(n1,a)*comb(n0,k-a) for a in range(lo,hi+1)]

def distribution(c,theta):
    vals=[F(v)*theta**i for i,v in enumerate(c)]; z=sum(vals)
    return [v/z for v in vals]

def mp_distribution(c,eta):
    # Shift the log weights, including very large |eta|.
    vals=[mp.log(v)+i*eta for i,v in enumerate(c)]
    largest=max(vals); vals=[mp.exp(v-largest) for v in vals]; z=sum(vals)
    return [v/z for v in vals]

def root(c,a,kind):
    def fn(eta):
        ps=mp_distribution(c,eta)
        if kind=='mle': return sum(i*p for i,p in enumerate(ps))-a
        if kind=='lower': return sum(ps[a:])-mp.mpf('.025')
        return sum(ps[:a+1])-mp.mpf('.025')
    lo=mp.mpf(-1); hi=mp.mpf(1)
    while fn(lo)*fn(hi)>0: lo*=2;hi*=2
    fl=fn(lo)
    for _ in range(360):
        mid=(lo+hi)/2;fm=fn(mid)
        if fl*fm<=0:hi=mid
        else:lo=mid;fl=fm
    return mp.exp((lo+hi)/2)

def bracket(c,a,kind,lo,hi):
    lps=distribution(c,F(lo));hps=distribution(c,F(hi))
    if kind=='lower':
        left=sum(lps[a:]);right=sum(hps[a:]);ok=left<F(1,40)<right
    elif kind=='upper':
        left=sum(lps[:a+1]);right=sum(hps[:a+1]);ok=left>F(1,40)>right
    else:
        left=sum(i*p for i,p in enumerate(lps));right=sum(i*p for i,p in enumerate(hps));ok=left<a<right
    check(ok,'rational_root_brackets')
    return {'lower':str(lo),'upper':str(hi),'value_at_lower':str(left),'value_at_upper':str(right)}

# Full-rank design certificates, including the additional-column transfer.
def signed(x,y,v): return [sum(a*b for a,b in zip(row,v))*(2*yi-1) for row,yi in zip(x,y)]
base=[[1,-1],[1,-1],[1,1],[1,1]];y=[0,1,0,1]
aug=[row+[yi] for row,yi in zip(base,y)]
check(det(matinfo(base,[1]*4))>0,'separation_certificates')
check(det(matinfo(aug,[1]*4))>0,'separation_certificates')
check(signed(aug,y,[F(-1,2),0,1])==[F(1,2)]*4,'separation_certificates')
quasi=[[1,-1],[1,0],[1,0],[1,1]]
check(signed(quasi,[0,0,1,1],[0,1])==[1,0,0,1],'separation_certificates')
for v in it.product(range(-8,9), repeat=2):
    if v==(0,0):continue
    check(min(signed(base,y,v))<0,'overlap_direction_grid')
    check(min(signed(quasi,[0,0,1,1],v))<=0,'quasi_no_strict_grid')

# Cauchy--Binet identity in exact arithmetic for nonconstant weights.
for d in [1,2,3]:
    for n in range(d,d+5):
        xs=[[F((i-2)**j) for j in range(d)] for i in range(n)]
        for scale in [F(1,3),F(1),F(2),F(5)]:
            odds=[scale**(i-2) for i in range(n)]
            ws=[o/(1+o)**2 for o in odds]
            lhs=det(matinfo(xs,ws));rhs=F(0)
            for s in it.combinations(range(n),d):
                p=F(1)
                for i in s:p*=ws[i]
                rhs+=det([xs[i] for i in s])**2*p
            check(lhs==rhs and lhs>0,'cauchy_binet')

# Analytic adjusted score vs direct differentiation, using arbitrary designs.
for xs in [base,aug,quasi,[[1,-2],[1,-1],[1,0],[1,2],[1,3]]]:
    d=len(xs[0]);n=len(xs);X=mp.matrix(xs)
    for ys in [([0,1]*(n//2+1))[:n],[0]*n,[1]*n]:
        yy=mp.matrix(ys)
        for v in [mp.mpf('-.7'),mp.mpf('0'),mp.mpf('.8')]:
            beta=mp.matrix([v+mp.mpf(j)/5 for j in range(d)])
            def obj(*bb):
                eta=X*mp.matrix(bb);pp=[1/(1+mp.exp(-z)) for z in eta]
                inf=X.T*mp.diag([q*(1-q) for q in pp])*X
                return sum(ys[i]*eta[i]-mp.log(1+mp.exp(eta[i])) for i in range(n))+mp.log(mp.det(inf))/2
            eta=X*beta;ps=[1/(1+mp.exp(-z)) for z in eta];ww=[p*(1-p) for p in ps]
            inf=X.T*mp.diag(ww)*X;inv=inf**-1
            hs=[ww[i]*(mp.matrix([xs[i]])*inv*mp.matrix(xs[i]))[0] for i in range(n)]
            close(sum(hs),mp.mpf(d),'hat_trace')
            for h in hs:check(h>=-mp.mpf('1e-75') and h<=1+mp.mpf('1e-75'),'hat_bounds')
            us=X.T*mp.matrix([ys[i]-ps[i]+hs[i]*(mp.mpf('.5')-ps[i]) for i in range(n)])
            for j in range(d):
                order=[0]*d;order[j]=1
                close(us[j],mp.diff(obj,tuple(beta),tuple(order)),'firth_score_derivatives')

# Saturated independent group fits for every pair of observed counts.
for n0 in range(1,8):
    for n1 in range(1,8):
        for s0 in range(n0+1):
            for s1 in range(n1+1):
                p0=F(2*s0+1,2*(n0+1));p1=F(2*s1+1,2*(n1+1))
                check(s0-n0*p0+F(1,2)-p0==0,'saturated_firth_scores')
                check(s1-n1*p1+F(1,2)-p1==0,'saturated_firth_scores')
                check(0<p0<1 and 0<p1<1,'saturated_finite_probabilities')

# General multidimensional subset DP agrees with exhaustive conditional moments.
for n in range(1,9):
    xs=[[(i%3)-1,((2*i+1)%5)-2] for i in range(n)]
    for base2,base3 in [(2,3),(3,2),(1,1)]:
        ws=[F(base2)**x[0]*F(base3)**x[1] for x in xs]
        for m in range(n+1):
            z,mean,cov=dp_moments(xs,ws,m)
            ez,em,ec,rows=enumerate_moments(xs,ws,m)
            check(z==ez,'subset_dp_normalizer')
            for i in range(2):
                check(mean[i]==em[i],'subset_dp_mean')
                for j in range(2):check(cov[i][j]==ec[i][j],'subset_dp_covariance')
            check(cov[0][0]>=0 and cov[1][1]>=0 and det(cov)>=0,'subset_covariance_psd')
            shifted=[[x[0]+7,x[1]-3] for x in xs]
            zz,mm,cc=dp_moments(shifted,ws,m)
            check(cc==cov and mm==[mean[0]+7*m,mean[1]-3*m],'within_stratum_shift')
            for alpha_odds in [F(1,3),F(2),F(5)]:
                ps=[alpha_odds*w/(1+alpha_odds*w) for w in ws]
                masses=[]
                for s,t,w in rows:
                    mass=F(1)
                    for i in range(n):mass*=ps[i] if i in s else 1-ps[i]
                    masses.append(mass)
                total=sum(masses)
                for (_,_,w),mass in zip(rows,masses):
                    check(mass/total==w/z,'conditional_intercept_cancellation')

# Exact tails and their finite-sample coverage over every possible observation.
thetas=[F(1,10),F(1,3),F(1,2),F(1),F(2),F(3),F(10)]
for n1 in range(1,7):
    for n0 in range(1,7):
        for k in range(n1+n0+1):
            low,c=table_coeff(n1,n0,k)
            for theta in thetas:
                ps=distribution(c,theta);m=len(ps)
                rt=[sum(ps[i:]) for i in range(m)];lt=[sum(ps[:i+1]) for i in range(m)]
                check(sum(ps)==1 and all(p>0 for p in ps),'conditional_mass')
                for i in range(m-1):
                    a=low+i;ratio=theta*F((n1-a)*(k-a),(a+1)*(n0-k+a+1))
                    check(ps[i+1]/ps[i]==ratio,'adjacent_mass_recurrence')
                for tail in [rt,lt]:
                    for t in set(tail):
                        check(sum(p for p,pv in zip(ps,tail) if pv<=t)<=t,'tail_superuniformity')
                for alpha in [F(1,20),F(1,10),F(1,5),F(1,2)]:
                    coverage=sum(p for i,p in enumerate(ps) if rt[i]>=alpha/2 and lt[i]>=alpha/2)
                    check(coverage>=1-alpha,'equal_tail_conditional_coverage')
                for p0 in [F(1,4),F(1,2),F(3,4)]:
                    p1=theta*p0/(1-p0+theta*p0)
                    masses=[F(comb(n1,a)*comb(n0,k-a))*p1**a*(1-p1)**(n1-a)*p0**(k-a)*(1-p0)**(n0-k+a)
                            for a in range(low,low+m)]
                    total=sum(masses)
                    check([v/total for v in masses]==ps,'binomial_to_noncentral_family')

# A full-rank original layer design can have zero conditional information.
original_layer_design=[[1,0,0],[1,0,0],[0,1,0],[0,1,1]]
check(det(matinfo(original_layer_design,[1]*4))>0,'conditional_information_loss')
check(dp_moments([[0],[0]],[F(1),F(1)],1)[2]==[[F(0)]] and dp_moments([[0],[1]],[F(1),F(2)],0)[2]==[[F(0)]],'conditional_information_loss')

# Terminal multidimensional example, including exact Newton direction.
xs=[[0,0],[1,0],[0,1],[1,1],[2,0]];ws=list(map(F,[1,2,3,6,4]))
z,mu,cov=dp_moments(xs,ws,2)
check(z==95 and mu==[F(184,95),F(99,95)],'capstone_vector_moments')
check(cov==[[F(7184,9025),F(-2256,9025)],[F(-2256,9025),F(3024,9025)]],'capstone_vector_moments')
check(det(cov)==F(9216,45125),'capstone_vector_moments')
uc=[F(-89,95),F(-4,95)];direction=[F(-305,192),F(-755,576)]
check([sum(cov[i][j]*direction[j] for j in range(2)) for i in range(2)]==uc,'capstone_newton_direction')
check(sum(uc[i]*direction[i] for i in range(2))>0,'capstone_newton_direction')

single=[1,9,9,1];second=[3,6,1];joint=conv(single,second)
check(joint==[3,33,82,66,15,1],'common_odds_convolution')
check(table_coeff(5,6,5)==(0,[6,75,200,150,30,1]),'common_odds_convolution')
check(joint!=table_coeff(5,6,5)[1],'common_odds_convolution')
for theta in thetas:
    check(conv(distribution(single,theta),distribution(second,theta))==distribution(joint,theta),'common_odds_convolution')

numeric={}
for name,c,a,bounds in [
    ('single_table',single,2,{'lower':('0.0674038648','0.0674038649'),'upper':('351.9974809912','351.9974809913'),'mle':('3.1054826165','3.1054826166')}),
    ('two_strata',joint,3,{'lower':('0.1561553005','0.1561553006'),'upper':('44.4883219661','44.4883219662'),'mle':('2.3748535384','2.3748535385')})]:
    numeric[name]={'coefficients':c,'observed':a}
    for kind in ['mle','lower','upper']:
        value=root(c,a,kind);lo,hi=bounds[kind]
        check(mp.mpf(lo)<value<mp.mpf(hi),'decimal_root_in_rational_bracket')
        numeric[name][kind]={'decimal_approx':mp.nstr(value,70),'exact_bracket':bracket(c,a,kind,lo,hi)}

zpair=mp.findroot(lambda z:z**3-z**2-z-3,2)
bpair=mp.log(zpair)
close(3-2/(1+mp.exp(-bpair))-2/(1+mp.exp(-2*bpair)),mp.mpf(0),'paired_finite_root')
numeric['three_pairs']={'exp_beta':mp.nstr(zpair,70),'beta':mp.nstr(bpair,70)}
z975=mp.sqrt(2)*mp.erfinv(mp.mpf('.95'))
waldmax=mp.log(3)+z975*4/mp.sqrt(3)
check(waldmax<6,'finite_wald_zero_coverage')
for y in [0,1]:
    estimate=(2*y-1)*mp.log(3)
    check(not(estimate-z975*4/mp.sqrt(3)<=6<=estimate+z975*4/mp.sqrt(3)),'finite_wald_zero_coverage')
numeric['wald']={'maximum_upper_endpoint':mp.nstr(waldmax,70),'true_beta':6,'coverage':0}

# A finite check of the proved fixed-interior second-order bias expansion.
bias=[]
for pp in ['0.2','0.35','0.5','0.65','0.8']:
    p=mp.mpf(pp);q=1-p
    for n in [50,100,200,400]:
        mass=q**n;expectation=mp.mpf(0)
        for s in range(n+1):
            expectation+=mass*mp.log((mp.mpf(s)+mp.mpf('.5'))/(n-s+mp.mpf('.5')))
            if s<n:mass*=mp.mpf(n-s)/(s+1)*p/q
        b=expectation-mp.log(p/q)
        check(abs(n*n*b)<5,'fixed_interior_bias_grid')
        bias.append({'p':pp,'n':n,'n_squared_bias':mp.nstr(n*n*b,30)})
result={'status':'PASS','arithmetic':'Fraction for finite certificates; mpmath 90 digits for displayed roots and derivatives',
        'assertions':sum(COUNTS.values()),'checks_by_group':COUNTS,'examples':numeric,
        'vector_layer':{'normalizer':str(z),'mean':[str(a) for a in mu],'covariance':[[str(a) for a in row] for row in cov],
                        'newton_direction':[str(a) for a in direction]},'fixed_interior_bias_grid':bias}
parser=argparse.ArgumentParser(description=__doc__);parser.add_argument('--output',type=Path);args=parser.parse_args()
text=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
if args.output:
    args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(text)
else:print(text,end='')
