Free-Loss Jacobian Spectral Target / stage2_bench.py
Failed on benchmark
1import sys, json, math, random
2from pathlib import Path
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))
10DEVICE='cuda' if torch.cuda.is_available() else 'cpu'
11
12def seed_all(s):
13 random.seed(s); np.random.seed(s); torch.manual_seed(s)
14 if torch.cuda.is_available():
15 try: torch.cuda.manual_seed_all(s)
16 except Exception: pass
17
18def free_moments(tau,K=2):
19 out=[]
20 for p in range(1,K+1):
21 n=p-1; a=p*tau; q=np.zeros(n+1); q[0]=1.; b=np.zeros(n+1)
22 for j in range(1,n+1): b[j]=a*((-1)**(j+1))
23 for k in range(1,n+1): q[k]=sum(j*b[j]*q[k-j] for j in range(1,k+1))/k
24 out.append(math.exp(-a)*sum(math.comb(p,j)*q[n-j] for j in range(n+1))/p)
25 return np.asarray(out,dtype=np.float32)
26
27def math_check():
28 rng=np.random.default_rng(11); n=28; tau=.5; L=round(tau*n); P=np.diag([1.]*(n-1)+[0.]); vals=[]
29 for _ in range(12):
30 B=np.eye(n,dtype=complex)
31 for _ in range(L):
32 z=(rng.normal(size=(n,n))+1j*rng.normal(size=(n,n)))/np.sqrt(2)
33 q,r=np.linalg.qr(z); d=np.diag(r); B=P@(q*(d/np.abs(d)).conj())@B
34 e=np.linalg.eigvalsh(B.conj().T@B).real; vals.append([e.mean(),(e**2).mean()])
35 emp=np.mean(vals,0); tar=free_moments(tau)
36 return {'tau':tau,'target':tar.tolist(),'empirical':emp.tolist(),'abs_error':np.abs(emp-tar).tolist(),'confirmed':bool(np.max(np.abs(emp-tar))<.03)}
37
38def baseline_one(seed,cfg):
39 seed_all(seed); d=get_dataset('tabular',seed,n_train=400,n_test=200)
40 _,metric,_=train_model(make_model('mlp_tiny',d['input_shape'],d['out_dim']),d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
41 return float(metric)
42
43def spectral_loss(net,x,target):
44 # Scalar-output Jacobian: one reverse-mode derivative gives per-example J rows.
45 x=x.detach().requires_grad_(True); y=net(x).reshape(-1)
46 g=torch.autograd.grad(y.sum(),x,create_graph=True)[0]
47 m1=g.pow(2).sum(1).mean(); m2=g.pow(4).sum(1).mean() / x.shape[1]
48 # Normalize moments as trace/N and trace((J^T J)^2)/N.
49 m=torch.stack([m1/x.shape[1],m2])
50 return ((torch.log(m+1e-5)-torch.log(target+1e-5))**2).sum(),m
51
52def idea_one(seed,cfg,capture=False):
53 seed_all(seed); d=get_dataset('tabular',seed,n_train=400,n_test=200)
54 net=make_model('mlp_tiny',d['input_shape'],d['out_dim'])
55 dev=torch.device(DEVICE); net=net.to(dev); x=d['xtr'].to(dev); y=d['ytr'].to(dev); xe=d['xte'].to(dev); ye=d['yte'].to(dev)
56 target=torch.tensor(free_moments(2/64),device=dev)
57 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); mse=nn.MSELoss()
58 for ep in range(cfg['epochs']):
59 net.train(); perm=torch.randperm(len(x),device=dev)
60 for st in range(0,len(x),128):
61 ix=perm[st:st+128]; pred=net(x[ix]); loss=mse(pred,y[ix])
62 # The intervention is applied periodically to a small probe batch only.
63 if st==0 and ep < max(1,cfg['epochs']//3):
64 sl,_=spectral_loss(net,x[ix[:8]],target); loss=loss+cfg['lam']*sl.clamp(max=10.)
65 opt.zero_grad(); loss.backward(); opt.step()
66 net.eval()
67 with torch.no_grad(): metric=float(mse(net(xe),ye).cpu())
68 if not capture:return metric
69 with torch.enable_grad(): sl,m=spectral_loss(net,xe[:16],target)
70 return {'metric':metric,'predicted_moments':target.cpu().tolist(),'observed_moments':m.detach().cpu().tolist(),'log_moment_rmse':float(torch.sqrt(sl/2).detach().cpu())}
71
72def main():
73 grid=[{'lr':lr,'epochs':8,'weight_decay':0.0,'lam':0.0} for lr in (.001,.003,.006)]
74 # sweep_baseline calls the supplied factory with each seed and preserves full seed results.
75 base=sweep_baseline(lambda c: (lambda seed: baseline_one(seed,c)),grid,seeds=(0,1,2,3))
76 best=base['best_cfg']; lrs=[c['lr'] for c in grid]
77 # Same lr union on idea side; three predetermined spectral weights.
78 idea_cfgs=[{'lr':lr,'epochs':8,'weight_decay':0.0,'lam':lam} for lr in lrs for lam in (.0001, .0003, .001)]
79 runs=[]
80 for c in idea_cfgs:
81 r=evaluate(lambda seed,c=c: idea_one(seed,c),seeds=SEEDS); runs.append((c,r))
82 ibest,ires=min(runs,key=lambda z:z[1]['mean'])
83 # Evaluate baseline best configuration on all eight paired seeds for make_report comparison.
84 base_full=evaluate(lambda seed: baseline_one(seed,best),seeds=SEEDS)
85 base['full']=base_full
86 sig=idea_one(0,ibest,True)
87 report=make_report('tabular','mlp_tiny',base,ires,extra={
88 'selected_idea_cfg':ibest,'math_check':math_check(),
89 'mechanism_signature':{'source':'trained mlp_tiny Jacobian on held-out tabular inputs','predicted_moments':sig['predicted_moments'],'observed_moments':sig['observed_moments'],'log_moment_rmse':sig['log_moment_rmse'],'confirmed':bool(sig['log_moment_rmse']<1.0)}})
90 report['idea_sweep']=[{'cfg':c,'result':r} for c,r in runs]
91 Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
92if __name__=='__main__': main()