Capacity-Shaped Binomial Bottleneck / bench_binomial.py
Beats tuned baseline
1import sys, json, math, random
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, sweep_baseline, make_report
7
8SEEDS=tuple(range(8)); N=8
9
10def beta_binomial(n=N):
11 a=[]
12 for y in range(n+1):
13 z=(math.lgamma(n+1)-math.lgamma(y+1)-math.lgamma(n-y+1)+math.lgamma(y+.5)+math.lgamma(n-y+.5)-math.lgamma(n+1)-math.log(math.pi))
14 a.append(math.exp(z))
15 a=np.asarray(a); return a/a.sum()
16QR_CPU=torch.tensor(beta_binomial(),dtype=torch.float32)
17
18def binomial_probs(x,n=N):
19 x=x.clamp(1e-5,1-1e-5); y=torch.arange(n+1,device=x.device,dtype=x.dtype)
20 lc=torch.lgamma(torch.tensor(float(n+1),device=x.device,dtype=x.dtype))-torch.lgamma(y+1)-torch.lgamma(torch.tensor(float(n),device=x.device,dtype=x.dtype)-y+1)
21 return torch.exp(lc+y*torch.log(x[...,None])+(n-y)*torch.log1p(-x[...,None]))
22
23def entropy(p):
24 z=p.clamp_min(1e-12); return -(z*z.log()).sum(-1)
25
26def bottleneck_terms(x):
27 p=binomial_probs(x); q=p.mean(0) # [latent, count]
28 mi=(entropy(q)-entropy(p).mean(0)).mean() # mean over latent coordinates
29 qr=QR_CPU.to(x.device)
30 qz=q.clamp_min(1e-12)
31 kl=(qz*(qz.log()-qr.log())).sum(-1).mean() # mean over latent coordinates
32 return mi,kl,p,q
33
34class BottleneckMLP(nn.Module):
35 def __init__(self,d,idea):
36 super().__init__(); self.idea=idea
37 self.enc=nn.Sequential(nn.Linear(d,64),nn.ReLU(),nn.Linear(64,8))
38 self.head=nn.Sequential(nn.ReLU(),nn.Linear(8,1))
39 def forward(self,x):
40 z=torch.sigmoid(self.enc(x))
41 if self.idea:
42 mi,kl,p,q=bottleneck_terms(z); return self.head(z),mi,kl,p,q,z
43 return self.head(z),None,None,None,None,z
44
45def seed_all(seed):
46 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
47 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
48
49def _run(seed,cfg,idea,device):
50 seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400)
51 xtr=ds['xtr'].float(); ytr=ds['ytr'].float().reshape(-1,1); xte=ds['xte'].float(); yte=ds['yte'].float().reshape(-1,1)
52 d=int(np.prod(tuple(xtr.shape[1:]))); model=BottleneckMLP(d,idea).to(device)
53 opt=torch.optim.Adam(model.parameters(),lr=float(cfg['lr']),weight_decay=float(cfg.get('weight_decay',0)))
54 xtr,ytr,xte,yte=[v.to(device) for v in (xtr,ytr,xte,yte)]
55 model.train()
56 for _ in range(int(cfg.get('epochs',30))):
57 for ix in torch.randperm(len(xtr),device=device).split(128):
58 out,mi,kl,_,_,_=model(xtr[ix]); loss=nn.functional.mse_loss(out,ytr[ix])
59 if idea: loss=loss+float(cfg['prior'])*kl-float(cfg['mi'])*mi
60 opt.zero_grad(); loss.backward(); opt.step()
61 model.eval()
62 with torch.no_grad():
63 pred,_,_,_,_,z=model(xte); metric=float(nn.functional.mse_loss(pred,yte).cpu())
64 mi,kl,p,q=bottleneck_terms(z); qbar=q.mean(0); endpoint=float((qbar[0]+qbar[-1]).cpu()); l1=float(torch.abs(qbar-QR_CPU.to(device)).sum().cpu())
65 return metric,{'predicted_prior_endpoint_mass':float(2*beta_binomial()[0]),'observed_endpoint_mass':endpoint,'predicted_prior_l1_to_reference':0.0,'observed_prior_l1_to_reference':l1,'observed_test_MI_nats':float(mi.cpu()),'observed_test_KL':float(kl.cpu()),'confirmed':abs(endpoint-2*beta_binomial()[0])<0.18}
66
67def run(seed,cfg,idea):
68 if torch.cuda.is_available():
69 try: return _run(seed,cfg,idea,torch.device('cuda'))
70 except Exception: pass
71 return _run(seed,cfg,idea,torch.device('cpu'))
72
73def main():
74 lrs=[1e-3,3e-3,1e-2]
75 grid=[{'lr':lr,'epochs':30,'weight_decay':wd,'prior':0.,'mi':0.} for lr in lrs for wd in (0.,1e-4)]
76 def make_base(cfg): return lambda seed: run(seed,cfg,False)[0]
77 base=sweep_baseline(make_base,grid,seeds=(0,1,2,3))
78 best=base.get('best_config',base.get('best_cfg',grid[0])); best=best if isinstance(best,dict) else grid[0]
79 bcfg={'lr':float(best['lr']),'epochs':30,'weight_decay':float(best.get('weight_decay',0.)),'prior':0.,'mi':0.}
80 bvals=[run(s,bcfg,False)[0] for s in SEEDS]
81 base_block={'best_config':bcfg,'sweep':base,'full':{'mean':float(np.mean(bvals)),'std':float(np.std(bvals,ddof=1)),'per_seed':bvals}}
82 idea_cfgs=[{'lr':lr,'epochs':30,'weight_decay':bcfg['weight_decay'],'prior':.04,'mi':.01} for lr in lrs]
83 results=[]
84 for cfg in idea_cfgs:
85 vals=[]; sigs=[]
86 for s in SEEDS:
87 v,sg=run(s,cfg,True); vals.append(v); sigs.append(sg)
88 results.append((float(np.mean(vals)),cfg,vals,sigs))
89 _,icfg,ivals,sigs=min(results,key=lambda x:x[0])
90 idea_res={'mean':float(np.mean(ivals)),'std':float(np.std(ivals,ddof=1)),'per_seed':ivals,'config':icfg}
91 sig={k:float(np.mean([s[k] for s in sigs])) for k in sigs[0] if k!='confirmed'}; sig['confirmed']=bool(np.mean([s['confirmed'] for s in sigs])>=.75)
92 rep=make_report('tabular','mlp_tiny',base_block,idea_res,{'mechanism_signature':sig,'audit':{'idea_configs':idea_cfgs,'baseline_lrs':lrs,'device':'cuda' if torch.cuda.is_available() else 'cpu'}})
93 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
94 print(json.dumps(rep,indent=2))
95if __name__=='__main__': main()