Gain-Weighted Cluster Co-Design / bench_experiment.py
Mechanism confirmed, baseline not beaten
1import sys, json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, make_report
7from bench.protocol import evaluate, sweep_baseline
8
9SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); NTR=400; NTE=200; EPOCHS=5; BATCH=128
10# CPU is selected intentionally: this workload is tiny and avoids shared-GPU fallback variability.
11DEVICE='cpu'
12
13def seed_all(s):
14 random.seed(s); np.random.seed(s); torch.manual_seed(s)
15
16def estimate_gains(net,x,device):
17 z=x[:32].detach().to(device).requires_grad_(True)
18 out=net(z).sum(); (gr,)=torch.autograd.grad(out,z)
19 a=gr.detach().abs().reshape(-1,8,3).mean((0,2)).cpu().numpy()+1e-8
20 return np.outer(a,1.0/a)
21
22def greedy_clusters(G,max_size=3):
23 cs=[{i} for i in range(8)]
24 while True:
25 best=None
26 for a in range(len(cs)):
27 for b in range(a):
28 if len(cs[a]|cs[b])>max_size: continue
29 score=max((math.log(float(G[i,j])+1e-8)+math.log(float(G[j,i])+1e-8)
30 for i in cs[a] for j in cs[b]),default=-1e99)
31 if best is None or score>best[0]: best=(score,a,b)
32 if best is None or best[0]<math.log(.01): break
33 _,a,b=best; m=cs[a]|cs[b]
34 cs=[c for k,c in enumerate(cs) if k not in (a,b)]+[m]
35 return sorted([sorted(c) for c in cs],key=lambda z:z[0])
36
37def baseline_run(seed,lr):
38 seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
39 net=make_model('rnn_small',d['input_shape'],d['out_dim'])
40 _,m,_=train_model(net,d,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
41 return float(m)
42
43def idea_run(seed,lr,return_sig=False):
44 seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
45 net=make_model('rnn_small',d['input_shape'],d['out_dim']).to(DEVICE)
46 x,y=d['xtr'].to(DEVICE),d['ytr'].to(DEVICE)
47 q=nn.Parameter(torch.zeros(8,device=DEVICE)); opt=torch.optim.Adam(list(net.parameters())+[q],lr=lr); lossf=nn.MSELoss()
48 for ep in range(EPOCHS):
49 net.eval(); G=estimate_gains(net,x,DEVICE); gt=torch.as_tensor(G,dtype=torch.float32)
50 dd=torch.exp(q-q.mean()); norm=dd.view(1,8,1)
51 P=greedy_clusters(G); owner={i:k for k,c in enumerate(P) for i in c}
52 cut=sum(G[i,j] for i in range(8) for j in range(8) if owner[i]!=owner[j])
53 perm=torch.randperm(len(x)); net.train()
54 for st in range(0,len(x),BATCH):
55 ix=perm[st:st+BATCH]
56 dd_batch=torch.exp(q-q.mean()); norm_batch=dd_batch.view(1,8,1)
57 mu_batch=torch.max((gt@dd_batch)/dd_batch)
58 pred=net(x[ix].view(-1,8,3)*norm_batch)
59 loss=lossf(pred,y[ix])+0.01*torch.nn.functional.softplus(mu_batch-1.0+0.02)+1e-4*cut
60 opt.zero_grad(); loss.backward(); opt.step()
61 net.eval(); xt=d['xte'].to(DEVICE)
62 with torch.no_grad(): metric=float(((net(xt)-d['yte'].to(DEVICE))**2).mean())
63 G2=estimate_gains(net,xt,DEVICE); dd2=torch.exp(q-q.mean()).detach().numpy()
64 sig={'mu':float(np.max((G2@dd2)/dd2)),'unscaled_mu':float(np.max(G2.sum(1))),'clusters':greedy_clusters(G2),'gain_mean':float(G2.mean())}
65 return (metric,sig) if return_sig else metric
66
67def main():
68 grid=[{'lr':v} for v in (1e-3,3e-3,6e-3)]
69 base=sweep_baseline(lambda c: lambda s: baseline_run(int(s),c['lr']),grid,seeds=SWEEP_SEEDS)
70 irs=[(c,evaluate(lambda s: idea_run(int(s),c['lr']),SWEEP_SEEDS)['mean']) for c in grid]
71 best_cfg=min(irs,key=lambda z:z[1])[0]; vals=[]; sigs=[]
72 for s in SEEDS:
73 v,sg=idea_run(s,best_cfg['lr'],True); vals.append(v); sigs.append(sg)
74 idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':8,'best_cfg':best_cfg,'sweep':[{'cfg':c,'mean':m} for c,m in irs]}
75 mm=float(np.mean([z['mu'] for z in sigs])); uu=float(np.mean([z['unscaled_mu'] for z in sigs]))
76 extra={'prediction':'trained-model gain weighting lowers normalized local certificate versus q=0','observed_normalized_mu':mm,'observed_unscaled_mu':uu,'relative_reduction':float((uu-mm)/max(uu,1e-12)),'clusters_from_trained_models':sigs[0]['clusters'],'confirmed':bool(mm<uu)}
77 print(json.dumps(make_report('dynamics','rnn_small',base,idea,extra),indent=2))
78if __name__=='__main__': main()