Gain-Weighted Cluster Co-Design / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()