Cholesky-Structured SPD Classifier / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, time
  2import numpy as np
  3import torch
  4import torch.nn.functional as F
  5
  6torch.set_default_dtype(torch.float64)
  7SEED=1145
  8np.random.seed(SEED); torch.manual_seed(SEED)
  9try:
 10    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 11except Exception:
 12    device=torch.device('cpu')
 13
 14def power(S,p):
 15    S=(S+S.transpose(-1,-2))/2
 16    w,V=torch.linalg.eigh(S)
 17    return (V*w.clamp_min(1e-14).pow(p).unsqueeze(-2))@V.transpose(-1,-2)
 18
 19def logspd(S): return power(S,0.0) if False else power_eigh_log(S)
 20def power_eigh_log(S):
 21    S=(S+S.transpose(-1,-2))/2; w,V=torch.linalg.eigh(S)
 22    return (V*torch.log(w.clamp_min(1e-14)).unsqueeze(-2))@V.transpose(-1,-2)
 23def low(X): return torch.tril(X,-1)
 24
 25def structured_logits(S,L,A,theta,delta=1e-8):
 26    d=S.shape[-1]; I=torch.eye(d,dtype=S.dtype,device=S.device)
 27    Sr=S+delta*I; K=torch.linalg.cholesky(Sr); Kp=power(Sr,theta/2)
 28    out=[]
 29    for k in range(L.shape[0]):
 30        P=L[k]@L[k].T; Lp=power(P,theta/2)
 31        t1=(low(K)-low(L[k]))*low(A[k])
 32        t2=(Kp-Lp)*A[k]
 33        out.append(t1.sum(dim=(-1,-2))+t2.sum(dim=(-1,-2))/(4*theta))
 34    return torch.stack(out,-1)
 35
 36def le_logits(S,protos):
 37    x=logspd(S); p=torch.stack([logspd(z) for z in protos])
 38    return -((x[:,None]-p[None])**2).sum((-1,-2))
 39
 40def synthetic(n=360,d=5,c=3):
 41    gs=[]; means=[]
 42    for k in range(c):
 43        q,_=torch.linalg.qr(torch.randn(d,d)); vals=torch.exp(torch.linspace(-.5,.6,d)+.3*k)
 44        means.append(q@torch.diag(vals)@q.T)
 45    for k in range(c):
 46        for _ in range(n//c):
 47            e=torch.randn(d,d); e=(e+e.T)/2
 48            gs.append((means[k]+.08*e@e.T, k))
 49    return torch.stack([x for x,y in gs]).to(device),torch.tensor([y for x,y in gs],device=device)
 50
 51def make_factors(protos):
 52    L=torch.linalg.cholesky(protos).detach().clone().requires_grad_()
 53    A=torch.tril(torch.randn_like(L)*.15,-1).detach().clone().requires_grad_()
 54    return L,A
 55
 56def mechanism_checks():
 57    d=5; X=torch.randn(d,d); S=X@X.T+1e-3*torch.eye(d)
 58    K=torch.linalg.cholesky(S); recon=(K@K.T-S).abs().max().item()
 59    # Prediction 1: triangular factors with positive diagonal always produce SPD matrices.
 60    deltas=[1e-10,1e-7,1e-4,1e-2]; mineig=[]
 61    for delta in deltas:
 62        Z=torch.randn(d,d); Z=torch.tril(Z)
 63        Z.diagonal().copy_(torch.nn.functional.softplus(Z.diagonal())+delta)
 64        mineig.append(float(torch.linalg.eigvalsh(Z@Z.T).min()))
 65    # Prediction 2: S^a = I + a log(S) + O(a^2), so the error is linear in theta.
 66    S2=torch.diag(torch.tensor([0.3,0.8,1.7,3.0,5.0]))
 67    logS=logspd(S2); thetas=[.4,.2,.1,.05,.025]
 68    power_errors=[]; scaled_errors=[]
 69    for th in thetas:
 70        err=power(S2,th/2)-torch.eye(d)
 71        power_errors.append(float(err.norm()))
 72        scaled_errors.append(float((err/(th/2)-logS).norm()))
 73    # Prediction 3: [K^(theta/2)-L^(theta/2)]/(4 theta)
 74    # tends to (log K-log L)/8 as theta -> 0 for diagonal commuting factors.
 75    K0=torch.diag(torch.tensor([.7,1.2,2.0,2.8,4.0])); L0=torch.diag(torch.tensor([1.1,.9,1.5,3.2,2.5]))
 76    A0=torch.tril(torch.ones(d,d),-1)+torch.eye(d)
 77    limit=((logspd(K0)-logspd(L0))*A0/8).sum().item()
 78    score_errors=[]
 79    for th in thetas:
 80        score=((power(K0,th/2)-power(L0,th/2))*A0/(4*th)).sum().item()
 81        score_errors.append(abs(score-limit))
 82    return {'cholesky_max_reconstruction_error':recon,
 83            'factor_deltas':deltas,'factor_min_eigenvalues':mineig,
 84            'power_thetas':thetas,'power_errors_vs_identity':power_errors,
 85            'power_log_limit_errors':scaled_errors,
 86            'scaled_power_limit_predicted':limit,'scaled_power_limit_errors':score_errors}
 87
 88def train_compare():
 89    S,y=synthetic(); c=3; protos=[]
 90    for k in range(c): protos.append(S[y==k].mean(0)+1e-3*torch.eye(S.shape[-1],device=device))
 91    protos=torch.stack(protos); L,A=make_factors(protos)
 92    opt=torch.optim.Adam([L,A],lr=.03); t0=time.time()
 93    for step in range(100):
 94        logits=structured_logits(S,L,A,1.0); loss=F.cross_entropy(logits,y)
 95        opt.zero_grad(); loss.backward(); opt.step()
 96    with torch.no_grad(): pred=structured_logits(S,L,A,1.0).argmax(1); acc=(pred==y).double().mean().item(); final=float(loss)
 97    structured_time=time.time()-t0
 98    t0=time.time(); le=le_logits(S,protos); le_acc=(le.argmax(1)==y).double().mean().item(); le_loss=float(F.cross_entropy(le,y)); le_time=time.time()-t0
 99    return {'structured_loss':final,'structured_accuracy':acc,'structured_seconds':structured_time,'logeuclidean_loss':le_loss,'logeuclidean_accuracy':le_acc,'logeuclidean_seconds':le_time}
100
101if __name__=='__main__':
102    try:
103        checks=mechanism_checks(); comparison=train_compare()
104        print(json.dumps({'device':str(device),'checks':checks,'comparison':comparison},indent=2))
105    except Exception as e:
106        if device.type=='cuda':
107            device=torch.device('cpu'); checks=mechanism_checks(); comparison=train_compare()
108            print(json.dumps({'device':'cpu','cuda_error':repr(e),'checks':checks,'comparison':comparison},indent=2))
109        else: raise