Separable Ky-Fan spectral regularization / bench_run.py
Mechanism confirmed, baseline not beaten
1from __future__ import annotations
2import json, random, sys
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10EPOCHS = 15
11BATCH = 128
12
13def kron_np(xs):
14 z = xs[0]
15 for x in xs[1:]: z = np.kron(z, x)
16 return z
17
18def verify_majorization(trials=40, n=3, d=3):
19 rng = np.random.default_rng(123)
20 failures, worst, ratios = 0, -float('inf'), []
21 for _ in range(trials):
22 aa, bb = [], []
23 for _ in range(n):
24 q = rng.normal(size=(d,d)); aa.append(q @ q.T + .05*np.eye(d))
25 q = rng.normal(size=(d,d)); bb.append(q @ q.T + .05*np.eye(d))
26 lm = np.linalg.eigvalsh(kron_np(aa)+kron_np(bb))[::-1]
27 av = [np.linalg.eigvalsh(x)[::-1] for x in aa]
28 bv = [np.linalg.eigvalsh(x)[::-1] for x in bb]
29 c = np.sort(kron_np(av)+kron_np(bv))[::-1]
30 gap = float(np.max(np.cumsum(lm)-np.cumsum(c)))
31 worst = max(worst, gap); failures += int(gap > 2e-8)
32 ratios.append(float(lm[0]/c[0]))
33 return {'trials':trials, 'failures':failures,
34 'worst_cumulative_gap':worst,
35 'max_eigenvalue_ratio_mean':float(np.mean(ratios))}
36
37def kron_t(xs):
38 z = xs[0]
39 for x in xs[1:]: z = torch.kron(z, x)
40 return z
41
42class PSDKron64(nn.Module):
43 """Shared MLP with its first 64x64 linear map represented as two PSD terms."""
44 def __init__(self, input_dim, out_dim, eps=1e-3):
45 super().__init__(); self.eps = eps
46 self.inp = nn.Linear(input_dim, 64)
47 self.LA = nn.ParameterList([nn.Parameter(.08*torch.randn(8,8)) for _ in range(2)])
48 self.LB = nn.ParameterList([nn.Parameter(.08*torch.randn(8,8)) for _ in range(2)])
49 self.bias = nn.Parameter(torch.zeros(64))
50 self.out = nn.Linear(64, out_dim)
51 def factors(self):
52 dev=self.LA[0].device; I=torch.eye(8,device=dev)
53 return ([x@x.T+self.eps*I for x in self.LA], [x@x.T+self.eps*I for x in self.LB])
54 def bound(self,k=1):
55 A,B=self.factors()
56 av=[torch.linalg.eigvalsh(x).flip(0) for x in A]
57 bv=[torch.linalg.eigvalsh(x).flip(0) for x in B]
58 c=torch.sort(kron_t(av)+kron_t(bv),descending=True).values
59 return c[:k].mean()
60 def exact_norm(self):
61 A,B=self.factors(); return torch.linalg.eigvalsh(kron_t(A)+kron_t(B))[-1]
62 def forward(self,x):
63 h=torch.relu(self.inp(x))
64 A,B=self.factors(); M=kron_t(A)+kron_t(B)
65 h=torch.relu(h@M.T+self.bias)
66 return self.out(h)
67
68def seed_all(seed):
69 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
70
71def baseline_one(cfg, seed, collect=False):
72 seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=200)
73 net=make_model('mlp_tiny',ds['input_shape'],ds['out_dim'])
74 net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['wd'],log=lambda _:None)
75 return float(metric)
76
77def idea_one(cfg, seed, collect=False):
78 seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=200)
79 dev='cuda' if torch.cuda.is_available() else 'cpu'
80 try:
81 net=PSDKron64(int(np.prod(ds['input_shape'])),ds['out_dim']).to(dev)
82 xtr,ytr=ds['xtr'].to(dev),ds['ytr'].to(dev); opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['wd'])
83 vals=[]; bounds=[]; norms=[]
84 for ep in range(EPOCHS):
85 net.train(); perm=torch.randperm(len(xtr),device=dev)
86 for i in range(0,len(xtr),BATCH):
87 ix=perm[i:i+BATCH]; pred=net(xtr[ix]); loss=((pred-ytr[ix])**2).mean()+cfg['lam']*net.bound(cfg['k'])
88 opt.zero_grad(); loss.backward(); opt.step(); vals.append(float(loss.detach().cpu()))
89 with torch.no_grad(): bounds.append(float(net.bound(cfg['k']).cpu())); norms.append(float(net.exact_norm().cpu()))
90 net.eval()
91 with torch.no_grad(): metric=float(((net(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean().cpu())
92 if collect:
93 b1=[]
94 with torch.no_grad():
95 for _ in norms: b1.append(float(net.bound(1).cpu()))
96 return metric, {'bound_R4':float(np.mean(bounds)),'bound_R1':float(np.mean(b1)),'exact_norm':float(np.mean(norms)),'ratio_R1':float(np.mean(np.array(norms)/np.array(b1)))}
97 return metric
98 except Exception:
99 if dev=='cuda':
100 torch.cuda.empty_cache(); torch.set_default_device('cpu'); return idea_one(cfg,seed,collect)
101 raise
102
103def main():
104 verification=verify_majorization()
105 # Union parity: every idea lr/wd is also evaluated by baseline.
106 grid=[{'lr':lr,'wd':wd,'lam':lam,'k':k} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4) for lam,k in ((0.0,1),)]
107 basegrid=[{'lr':x['lr'],'wd':x['wd']} for x in grid]
108 base=sweep_baseline(lambda c: (lambda s: baseline_one(c,s)),basegrid,seeds=tuple(range(4)))
109 # Baseline sweep includes all idea learning rates; final best is evaluated on 8 seeds.
110 idea_grid=[{'lr':lr,'wd':wd,'lam':lam,'k':k} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4) for lam,k in ((0.001,1),(0.003,1),(0.001,4))]
111 tried=[]
112 for c in idea_grid:
113 r=evaluate(lambda s: idea_one(c,s),seeds=tuple(range(4))); tried.append({'cfg':c,'mean':r['mean']})
114 best=min(tried,key=lambda x:x['mean'])['cfg']
115 idea=evaluate(lambda s: idea_one(best,s),seeds=SEEDS)
116 # Trained-model signature, independently measured on two paired trained models.
117 sig=[]
118 for s in (0,1):
119 _,z=idea_one(best,s,True); sig.append(z)
120 signature={'quantity':'exact operator norm / separable R1 bound on trained models','predicted':'ratio <= 1 by the Ky-Fan upper bound','observed_mean_ratio':float(np.mean([z['ratio_R1'] for z in sig])),'observed':sig,'confirmed':all(z['ratio_R1']<=1.0001 for z in sig)}
121 base['idea_union_sweep']=tried
122 rep=make_report('tabular','mlp_tiny',base,idea,{'mechanism_signature':signature,'verification':verification,'idea_best_cfg':best})
123 rep['stage2_note']='Baseline and idea use same tabular data, epochs, batch size, Adam, learning-rate/weight-decay union; idea changes only the first hidden operator and adds Ky-Fan penalty.'
124 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
125 print(json.dumps(rep,indent=2))
126if __name__=='__main__': main()