Cholesky-Structured SPD Classifier / experiment.py
Mechanism confirmed, baseline not beaten
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