import sys, json, math, random import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, sweep_baseline, make_report SEEDS=tuple(range(8)); N=8 def beta_binomial(n=N): a=[] for y in range(n+1): 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)) a.append(math.exp(z)) a=np.asarray(a); return a/a.sum() QR_CPU=torch.tensor(beta_binomial(),dtype=torch.float32) def binomial_probs(x,n=N): x=x.clamp(1e-5,1-1e-5); y=torch.arange(n+1,device=x.device,dtype=x.dtype) 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) return torch.exp(lc+y*torch.log(x[...,None])+(n-y)*torch.log1p(-x[...,None])) def entropy(p): z=p.clamp_min(1e-12); return -(z*z.log()).sum(-1) def bottleneck_terms(x): p=binomial_probs(x); q=p.mean(0) # [latent, count] mi=(entropy(q)-entropy(p).mean(0)).mean() # mean over latent coordinates qr=QR_CPU.to(x.device) qz=q.clamp_min(1e-12) kl=(qz*(qz.log()-qr.log())).sum(-1).mean() # mean over latent coordinates return mi,kl,p,q class BottleneckMLP(nn.Module): def __init__(self,d,idea): super().__init__(); self.idea=idea self.enc=nn.Sequential(nn.Linear(d,64),nn.ReLU(),nn.Linear(64,8)) self.head=nn.Sequential(nn.ReLU(),nn.Linear(8,1)) def forward(self,x): z=torch.sigmoid(self.enc(x)) if self.idea: mi,kl,p,q=bottleneck_terms(z); return self.head(z),mi,kl,p,q,z return self.head(z),None,None,None,None,z def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def _run(seed,cfg,idea,device): seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400) xtr=ds['xtr'].float(); ytr=ds['ytr'].float().reshape(-1,1); xte=ds['xte'].float(); yte=ds['yte'].float().reshape(-1,1) d=int(np.prod(tuple(xtr.shape[1:]))); model=BottleneckMLP(d,idea).to(device) opt=torch.optim.Adam(model.parameters(),lr=float(cfg['lr']),weight_decay=float(cfg.get('weight_decay',0))) xtr,ytr,xte,yte=[v.to(device) for v in (xtr,ytr,xte,yte)] model.train() for _ in range(int(cfg.get('epochs',30))): for ix in torch.randperm(len(xtr),device=device).split(128): out,mi,kl,_,_,_=model(xtr[ix]); loss=nn.functional.mse_loss(out,ytr[ix]) if idea: loss=loss+float(cfg['prior'])*kl-float(cfg['mi'])*mi opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): pred,_,_,_,_,z=model(xte); metric=float(nn.functional.mse_loss(pred,yte).cpu()) 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()) 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} def run(seed,cfg,idea): if torch.cuda.is_available(): try: return _run(seed,cfg,idea,torch.device('cuda')) except Exception: pass return _run(seed,cfg,idea,torch.device('cpu')) def main(): lrs=[1e-3,3e-3,1e-2] grid=[{'lr':lr,'epochs':30,'weight_decay':wd,'prior':0.,'mi':0.} for lr in lrs for wd in (0.,1e-4)] def make_base(cfg): return lambda seed: run(seed,cfg,False)[0] base=sweep_baseline(make_base,grid,seeds=(0,1,2,3)) best=base.get('best_config',base.get('best_cfg',grid[0])); best=best if isinstance(best,dict) else grid[0] bcfg={'lr':float(best['lr']),'epochs':30,'weight_decay':float(best.get('weight_decay',0.)),'prior':0.,'mi':0.} bvals=[run(s,bcfg,False)[0] for s in SEEDS] base_block={'best_config':bcfg,'sweep':base,'full':{'mean':float(np.mean(bvals)),'std':float(np.std(bvals,ddof=1)),'per_seed':bvals}} idea_cfgs=[{'lr':lr,'epochs':30,'weight_decay':bcfg['weight_decay'],'prior':.04,'mi':.01} for lr in lrs] results=[] for cfg in idea_cfgs: vals=[]; sigs=[] for s in SEEDS: v,sg=run(s,cfg,True); vals.append(v); sigs.append(sg) results.append((float(np.mean(vals)),cfg,vals,sigs)) _,icfg,ivals,sigs=min(results,key=lambda x:x[0]) idea_res={'mean':float(np.mean(ivals)),'std':float(np.std(ivals,ddof=1)),'per_seed':ivals,'config':icfg} 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) 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'}}) with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()