#!/usr/bin/env python3
"""Exact, standard-library checks for the spline-smoothing learning task.
Run with Python 3; no network, packages, or write access are needed.
All displayed rationals are generated from Fraction arithmetic.
"""
from fractions import Fraction as F
from itertools import combinations
import json

def transpose(A): return [list(x) for x in zip(*A)]
def mul(A,B): return [[sum(a*b for a,b in zip(row,col)) for col in zip(*B)] for row in A]
def mv(A,v): return [sum(a*b for a,b in zip(row,v)) for row in A]
def dot(a,b): return sum(x*y for x,y in zip(a,b))
def eye(n): return [[F(i==j) for j in range(n)] for i in range(n)]
def solve(A,b):
    n=len(A); T=[[F(x) for x in row]+[F(bi)] for row,bi in zip(A,b)]
    for j in range(n):
        k=next((k for k in range(j,n) if T[k][j]),None)
        if k is None: raise ValueError('singular system')
        T[j],T[k]=T[k],T[j]; d=T[j][j]; T[j]=[x/d for x in T[j]]
        for i in range(n):
            if i!=j:
                d=T[i][j];T[i]=[a-d*b for a,b in zip(T[i],T[j])]
    return [r[-1] for r in T]
def inverse(A): return transpose([solve(A,col) for col in eye(len(A))])
def det(A):
    if not A:return F(1)
    return sum((-1)**j*A[0][j]*det([r[:j]+r[j+1:] for r in A[1:]]) for j in range(len(A)))
def trace(A):return sum(A[i][i] for i in range(len(A)))
def variation(v):
    signs=[1 if x>0 else -1 for x in v if x]
    return sum(a!=b for a,b in zip(signs,signs[1:]))
def basis(t,p,x):
    N=len(t)-p-1
    if x==t[-1]:return [F(i==N-1) for i in range(N)]
    v=[F(t[i]<=x<t[i+1]) for i in range(len(t)-1)]
    for q in range(1,p+1):
        out=[]
        for i in range(len(v)-1):
            a=(x-t[i])*v[i]/(t[i+q]-t[i]) if t[i+q]!=t[i] else 0
            b=(t[i+q+1]-x)*v[i+1]/(t[i+q+1]-t[i+1]) if t[i+q+1]!=t[i+1] else 0
            out.append(a+b)
        v=out
    return v

def derivative(t,p,c):
    """Derivative in the trimmed clamped basis; positive spans only here."""
    assert p>=1 and len(c)==len(t)-p-1
    d=[]
    for i in range(1,len(c)):
        span=t[i+p]-t[i]
        if span<=0: raise ValueError('split full-multiplicity discontinuities first')
        d.append(p*(c[i]-c[i-1])/span)
    return t[1:-1],p-1,d

def main():
    v=list(map(F,[1,-2,1])); lam=F(1,9);y=list(map(F,[0,1,0]))
    A=[[F(3,2)*x*z for z in v] for x in v]
    H=[[F(i==j)+lam*A[i][j] for j in range(3)] for i in range(3)]
    S=inverse(H);z=mv(S,y);e=[yi-zi for yi,zi in zip(y,z)]
    energy=dot(z,mv(A,z));rss=dot(e,e);edf=trace(S);vtrace=trace(mul(S,S))
    loo=[e[i]/(1-S[i][i]) for i in range(3)]
    cv=dot(loo,loo)/3;gcv=(rss/3)/(1-edf/3)**2
    assert z==[F(1,6),F(2,3),F(1,6)]
    assert (energy,rss,edf,vtrace,cv,gcv)==(F(3,2),F(1,6),F(5,2),F(9,4),F(3),F(2))
    # Three-point flat LOO and GCV persist over distinct positive penalties.
    flat=[]
    for l in [F(1,100),F(1,9),F(1),F(100)]:
        s=inverse([[F(i==j)+l*A[i][j] for j in range(3)] for i in range(3)])
        f=mv(s,y);r=[yi-fi for yi,fi in zip(y,f)];rr=[r[i]/(1-s[i][i]) for i in range(3)]
        a=dot(rr,rr)/3;b=(dot(r,r)/3)/(1-trace(s)/3)**2
        assert (a,b)==(F(3),F(2));flat.append([l,a,b])
    # The anchored kernel plus affine coefficients reproduces all fitted values.
    X=[F(0),F(1),F(2)]
    K=[[min(x,t)**2*(3*max(x,t)-min(x,t))/6 for t in X] for x in X]
    c=[F(-3,2),F(3),F(-3,2)];beta=[F(1,6),F(3,4)]
    assert [beta[0]+beta[1]*x+k for x,k in zip(X,mv(K,c))]==z
    assert sum(c)==0 and dot(c,X)==0
    # All minors of a five-by-four ordinary-value collocation matrix.
    t=list(map(F,[0,0,0,1,2,2,2]));queries=[F(0),F(1,2),F(1),F(3,2),F(2)]
    C=[basis(t,2,x) for x in queries];nminors=0;zeros=0
    for k in range(1,5):
        for ii in combinations(range(5),k):
            for jj in combinations(range(4),k):
                d=det([[C[i][j] for j in jj] for i in ii]);assert d>=0
                nminors+=1;zeros+=d==0
    coeff=list(map(F,[1,-2,2,-1]));values=mv(C,coeff)
    assert values==list(map(F,[1,F(-3,4),0,F(3,4),-1]))
    assert variation(coeff)==variation(values)==3
    # Continuous curve and affine lines reproduce at every rational test point.
    greville=[F(0),F(1,2),F(3,2),F(2)]
    for j in range(41):
        x=F(j,20);vals=basis(t,2,x)
        p=1-6*x+5*x*x if x<=1 else 4*(x-1)-5*(x-1)**2
        assert dot(vals,coeff)==p and dot(vals,greville)==x
    # Nonuniform knots: discrete and integrated curvature nullspaces differ.
    linear=[F(0),F(1,2),F(2),F(3)]
    dif=[linear[i]-2*linear[i+1]+linear[i+2] for i in range(2)]
    assert dot(dif,dif)==F(5,4)
    nt=list(map(F,[0,0,0,1,3,3,3])); nc=list(map(F,[0,1,2,3]))
    dt,dp,dc=derivative(nt,2,nc)
    ddt,ddp,ddc=derivative(dt,dp,dc)
    assert ddp==0 and ddc==[F(-4,3),F(1,6)]
    wrongnull_energy=sum((ddt[i+1]-ddt[i])*q*q for i,q in enumerate(ddc))
    assert wrongnull_energy==F(11,6)
    at,ap,ac=derivative(nt,2,linear)
    _,_,aac=derivative(at,ap,ac)
    assert ac==[F(1)]*3 and aac==[F(0)]*2
    # Correlated additive spline components with true curvature penalties.
    u=list(map(F,[1,-2,1]));w=list(map(F,[-1,-1,2]));yy=[a+2*b for a,b in zip(u,w)]
    HH=[[F(12),F(3)],[F(3),F(12)]];bb=[dot(u,yy),dot(w,yy)];theta=solve(HH,bb)
    assert theta==[F(11,15),F(16,15)]
    def obj(a,b):return sum((q-a*x-b*z)**2 for q,x,z in zip(yy,u,w))+6*(a*a+b*b)
    path=[];a=b=F(0)
    for k in range(5):
        a=1-b/4;b=F(5,4)-a/4;path.append([k+1,a,b,obj(a,b)])
    assert path[0][1:]==[F(1),F(1),F(18)]
    assert path[1][1:]==[F(3,4),F(17,16),F(1101,64)]
    assert obj(*theta)==F(86,5)
    for k in range(1,5):assert path[k][2]-theta[1]==(path[k-1][2]-theta[1])/16
    B=transpose([u,w]);Ss=mul(mul(B,inverse(HH)),transpose(B))
    Ss=[[q+F(1,3) for q in row] for row in Ss]
    assert (trace(Ss),trace(mul(Ss,Ss)))==(F(29,15),F(331,225))
    result={'smoothing':{'fitted':z,'roughness':energy,'rss':rss,'objective':rss+lam*energy,'edf':edf,'variance_trace':vtrace,'loo_residuals':loo,'cv':cv,'gcv':gcv,'flat_scores':flat},'variation':{'collocation':C,'all_minors_checked':nminors,'zero_minors':zeros,'values':values,'strong_sign_changes':variation(values)},'penalty_mismatch':{'affine_curve_difference_penalty':dot(dif,dif),'index_affine_derivative_coefficients':dc,'index_affine_second_derivative_coefficients':ddc,'index_affine_curve_roughness':wrongnull_energy},'additive':{'coefficients':theta,'objective':obj(*theta),'sweeps':path,'edf':trace(Ss),'variance_trace':trace(mul(Ss,Ss))}}
    print(json.dumps(result,ensure_ascii=False,indent=2,default=lambda x:str(x)))

if __name__=='__main__':main()
