Capacity-Shaped Binomial Bottleneck / bench_binomial.py

✓✓ Beats tuned baseline

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