Separable Ky-Fan spectral regularization / experiment.py
Mechanism confirmed, baseline not beaten
1import json, time, math, random
2import numpy as np
3import torch
4
5SEED = 1729
6random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
7torch.set_num_threads(4)
8
9def kron_all(xs):
10 z = xs[0]
11 for x in xs[1:]: z = torch.kron(z, x)
12 return z
13
14def numpy_kron_all(xs):
15 z = xs[0]
16 for x in xs[1:]: z = np.kron(z, x)
17 return z
18
19def verify_majorization(trials=100, n=3, d=3):
20 worst = -np.inf; failures = 0; max_k_gap = np.zeros(numpy_kron_all([np.eye(d)]*n).shape[0])
21 tight_ratios=[]
22 for t in range(trials):
23 As=[]; Bs=[]
24 for i in range(n):
25 x=np.random.randn(d,d); As.append(x@x.T + .05*np.eye(d))
26 x=np.random.randn(d,d); Bs.append(x@x.T + .05*np.eye(d))
27 M=numpy_kron_all(As)+numpy_kron_all(Bs)
28 lm=np.linalg.eigvalsh(M)[::-1]
29 av=[np.linalg.eigvalsh(x)[::-1] for x in As]
30 bv=[np.linalg.eigvalsh(x)[::-1] for x in Bs]
31 c=numpy_kron_all(av)+numpy_kron_all(bv); c=np.sort(c)[::-1]
32 gaps=np.cumsum(lm)-np.cumsum(c)
33 worst=max(worst, float(gaps.max()))
34 max_k_gap=np.maximum(max_k_gap, gaps)
35 failures += int(np.any(gaps > 2e-9))
36 tight_ratios.append(lm[0]/c[0])
37 return {'trials':trials, 'failures':failures, 'worst_cumulative_gap':worst,
38 'max_gap_by_k':max_k_gap.tolist(), 'max_eigenvalue_ratio_mean':float(np.mean(tight_ratios)),
39 'max_eigenvalue_ratio_min':float(np.min(tight_ratios)), 'max_eigenvalue_ratio_max':float(np.max(tight_ratios))}
40
41class SepPSD(torch.nn.Module):
42 def __init__(self, n=3, d=3, eps=1e-3):
43 super().__init__(); self.n=n; self.d=d; self.eps=eps
44 self.LA=torch.nn.ParameterList([torch.nn.Parameter(.15*torch.randn(d,d)) for _ in range(n)])
45 self.LB=torch.nn.ParameterList([torch.nn.Parameter(.15*torch.randn(d,d)) for _ in range(n)])
46 def factors(self):
47 I=torch.eye(self.d, device=self.LA[0].device)
48 A=[L@L.T+self.eps*I for L in self.LA]
49 B=[L@L.T+self.eps*I for L in self.LB]
50 return A,B
51 def bound(self, k=1):
52 A,B=self.factors()
53 ae=[torch.linalg.eigvalsh(x).flip(0) for x in A]
54 be=[torch.linalg.eigvalsh(x).flip(0) for x in B]
55 c=kron_all(ae)+kron_all(be)
56 return torch.sort(c,descending=True).values[:k].sum()/k
57 def matrix(self):
58 A,B=self.factors(); return kron_all(A)+kron_all(B)
59 def forward(self,x, control=None):
60 M=self.matrix()
61 if control is not None:
62 r=self.bound(1)
63 scale=torch.clamp(torch.as_tensor(control,device=x.device)/r, max=1.0)
64 M=M*scale
65 return x@M.T
66
67def train(mode, seed=1729, steps=250, lr=.12):
68 torch.manual_seed(seed)
69 dev='cuda' if torch.cuda.is_available() else 'cpu'
70 try:
71 model=SepPSD().to(dev)
72 g=torch.Generator(device=dev); g.manual_seed(seed+1)
73 X=torch.randn(512,27,device=dev,generator=g)
74 # A fixed, well-conditioned PSD target gives a nontrivial operator fitting task.
75 q,_=torch.linalg.qr(torch.randn(27,27,device=dev,generator=g))
76 target=q@torch.diag(torch.linspace(.15,1.0,27,device=dev))@q.T
77 Y=X@target.T
78 opt=torch.optim.Adam(model.parameters(),lr=lr)
79 losses=[]; norms=[]; bounds=[]; spikes=0
80 t0=time.perf_counter()
81 for step in range(steps):
82 ix=torch.arange((step*64)%448,(step*64)%448+64,device=dev)
83 xb,yb=X[ix],Y[ix]
84 pred=model(xb, control=2.0 if mode=='controller' else None)
85 loss=((pred-yb)**2).mean()
86 if mode=='penalty': loss=loss + .01*model.bound(1)
87 if not torch.isfinite(loss): spikes += 1; break
88 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1000.0); opt.step()
89 losses.append(float(loss.detach().cpu()))
90 with torch.no_grad():
91 M=model.matrix(); norms.append(float(torch.linalg.eigvalsh(M)[-1].cpu())); bounds.append(float(model.bound(1).cpu()))
92 elapsed=time.perf_counter()-t0
93 final_loss=losses[-1] if losses else float('inf')
94 return {'mode':mode,'device':dev,'final_loss':final_loss,'best_loss':min(losses) if losses else float('inf'),
95 'steps_completed':len(losses),'nonfinite_steps':spikes,'max_exact_norm':max(norms) if norms else None,
96 'final_exact_norm':norms[-1] if norms else None,'final_R1':bounds[-1] if bounds else None,
97 'norm_over_R1':norms[-1]/bounds[-1] if bounds else None,'seconds':elapsed}
98 except Exception as e:
99 if dev=='cuda':
100 torch.cuda.empty_cache()
101 # retry CPU by temporarily hiding CUDA is awkward; report error and caller reruns subprocess-free CPU path
102 return {'mode':mode,'error':repr(e)}
103
104def main():
105 verification=verify_majorization()
106 results=[]
107 for mode in ['baseline','penalty','controller']:
108 r=train(mode)
109 if 'error' in r and r.get('device')=='cuda':
110 # CPU fallback in-process
111 old=torch.cuda.is_available
112 torch.cuda.is_available=lambda: False
113 r=train(mode)
114 torch.cuda.is_available=old
115 results.append(r)
116 out={'seed':SEED,'verification':verification,'training':results}
117 with open('results.json','w') as f: json.dump(out,f,indent=2)
118 print(json.dumps(out,indent=2))
119if __name__=='__main__': main()