import json, time import numpy as np import torch import torch.nn.functional as F torch.set_default_dtype(torch.float64) SEED=1145 np.random.seed(SEED); torch.manual_seed(SEED) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') except Exception: device=torch.device('cpu') def power(S,p): S=(S+S.transpose(-1,-2))/2 w,V=torch.linalg.eigh(S) return (V*w.clamp_min(1e-14).pow(p).unsqueeze(-2))@V.transpose(-1,-2) def logspd(S): return power(S,0.0) if False else power_eigh_log(S) def power_eigh_log(S): S=(S+S.transpose(-1,-2))/2; w,V=torch.linalg.eigh(S) return (V*torch.log(w.clamp_min(1e-14)).unsqueeze(-2))@V.transpose(-1,-2) def low(X): return torch.tril(X,-1) def structured_logits(S,L,A,theta,delta=1e-8): d=S.shape[-1]; I=torch.eye(d,dtype=S.dtype,device=S.device) Sr=S+delta*I; K=torch.linalg.cholesky(Sr); Kp=power(Sr,theta/2) out=[] for k in range(L.shape[0]): P=L[k]@L[k].T; Lp=power(P,theta/2) t1=(low(K)-low(L[k]))*low(A[k]) t2=(Kp-Lp)*A[k] out.append(t1.sum(dim=(-1,-2))+t2.sum(dim=(-1,-2))/(4*theta)) return torch.stack(out,-1) def le_logits(S,protos): x=logspd(S); p=torch.stack([logspd(z) for z in protos]) return -((x[:,None]-p[None])**2).sum((-1,-2)) def synthetic(n=360,d=5,c=3): gs=[]; means=[] for k in range(c): q,_=torch.linalg.qr(torch.randn(d,d)); vals=torch.exp(torch.linspace(-.5,.6,d)+.3*k) means.append(q@torch.diag(vals)@q.T) for k in range(c): for _ in range(n//c): e=torch.randn(d,d); e=(e+e.T)/2 gs.append((means[k]+.08*e@e.T, k)) return torch.stack([x for x,y in gs]).to(device),torch.tensor([y for x,y in gs],device=device) def make_factors(protos): L=torch.linalg.cholesky(protos).detach().clone().requires_grad_() A=torch.tril(torch.randn_like(L)*.15,-1).detach().clone().requires_grad_() return L,A def mechanism_checks(): d=5; X=torch.randn(d,d); S=X@X.T+1e-3*torch.eye(d) K=torch.linalg.cholesky(S); recon=(K@K.T-S).abs().max().item() # Prediction 1: triangular factors with positive diagonal always produce SPD matrices. deltas=[1e-10,1e-7,1e-4,1e-2]; mineig=[] for delta in deltas: Z=torch.randn(d,d); Z=torch.tril(Z) Z.diagonal().copy_(torch.nn.functional.softplus(Z.diagonal())+delta) mineig.append(float(torch.linalg.eigvalsh(Z@Z.T).min())) # Prediction 2: S^a = I + a log(S) + O(a^2), so the error is linear in theta. S2=torch.diag(torch.tensor([0.3,0.8,1.7,3.0,5.0])) logS=logspd(S2); thetas=[.4,.2,.1,.05,.025] power_errors=[]; scaled_errors=[] for th in thetas: err=power(S2,th/2)-torch.eye(d) power_errors.append(float(err.norm())) scaled_errors.append(float((err/(th/2)-logS).norm())) # Prediction 3: [K^(theta/2)-L^(theta/2)]/(4 theta) # tends to (log K-log L)/8 as theta -> 0 for diagonal commuting factors. 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])) A0=torch.tril(torch.ones(d,d),-1)+torch.eye(d) limit=((logspd(K0)-logspd(L0))*A0/8).sum().item() score_errors=[] for th in thetas: score=((power(K0,th/2)-power(L0,th/2))*A0/(4*th)).sum().item() score_errors.append(abs(score-limit)) return {'cholesky_max_reconstruction_error':recon, 'factor_deltas':deltas,'factor_min_eigenvalues':mineig, 'power_thetas':thetas,'power_errors_vs_identity':power_errors, 'power_log_limit_errors':scaled_errors, 'scaled_power_limit_predicted':limit,'scaled_power_limit_errors':score_errors} def train_compare(): S,y=synthetic(); c=3; protos=[] for k in range(c): protos.append(S[y==k].mean(0)+1e-3*torch.eye(S.shape[-1],device=device)) protos=torch.stack(protos); L,A=make_factors(protos) opt=torch.optim.Adam([L,A],lr=.03); t0=time.time() for step in range(100): logits=structured_logits(S,L,A,1.0); loss=F.cross_entropy(logits,y) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): pred=structured_logits(S,L,A,1.0).argmax(1); acc=(pred==y).double().mean().item(); final=float(loss) structured_time=time.time()-t0 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 return {'structured_loss':final,'structured_accuracy':acc,'structured_seconds':structured_time,'logeuclidean_loss':le_loss,'logeuclidean_accuracy':le_acc,'logeuclidean_seconds':le_time} if __name__=='__main__': try: checks=mechanism_checks(); comparison=train_compare() print(json.dumps({'device':str(device),'checks':checks,'comparison':comparison},indent=2)) except Exception as e: if device.type=='cuda': device=torch.device('cpu'); checks=mechanism_checks(); comparison=train_compare() print(json.dumps({'device':'cpu','cuda_error':repr(e),'checks':checks,'comparison':comparison},indent=2)) else: raise