import json, random, sys import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3); EPOCHS=12 GRID=[{'lr':1e-3,'epochs':EPOCHS},{'lr':3e-3,'epochs':EPOCHS},{'lr':1e-2,'epochs':EPOCHS}] def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def lunar_loss(z,c,tau=0.08): a,b=z[c==0],z[c==1] if len(a)==0 or len(b)==0: return z.sum()*0 d=torch.cdist(a,b) sa=-tau*torch.logsumexp(-d/tau,dim=1) sb=-tau*torch.logsumexp(-d/tau,dim=0) diam=torch.pdist(z).max().clamp_min(1e-5) if len(z)>1 else z.new_tensor(1.) return (sa.mean()+sb.mean())/(2*diam) class LunarModel(nn.Module): def __init__(self,input_shape,out_dim): super().__init__(); d=int(np.prod(input_shape)) self.backbone=nn.Sequential(nn.Linear(d,64),nn.ReLU(),nn.Linear(64,2)) self.head=nn.Linear(2,out_dim) def forward(self,x): z=self.backbone(x.flatten(1)); return z,self.head(z) class OutputWrapper(nn.Module): def __init__(self,m): super().__init__(); self.m=m def forward(self,x): return self.m(x)[1] def run(cfg,seed,idea=False,return_signature=False): seed_all(seed); d=get_dataset('tabular',seed,n_train=400,n_test=200) model=LunarModel(d['input_shape'],d['out_dim']) if not idea: _,metric,_=train_model(OutputWrapper(model),d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *a,**k:None) return float(metric) opt=torch.optim.Adam(model.parameters(),lr=cfg['lr']); x,y=d['xtr'],d['ytr']; colors=torch.arange(len(x))%2 for _ in range(cfg['epochs']): for ix in torch.randperm(len(x)).split(128): z,p=model(x[ix]); loss=((p-y[ix])**2).mean()+0.05*lunar_loss(z,colors[ix]) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): z,p=model(d['xte']); mse=((p-d['yte'])**2).mean().item() zz,_=model(d['xtr']); cross=torch.cdist(zz[colors==0],zz[colors==1]).min(1).values.mean().item() return (float(mse),float(cross)) if return_signature else float(mse) def base_factory(cfg): return lambda seed: run(cfg,seed,False) def idea_factory(cfg): return lambda seed: run(cfg,seed,True) def main(): base=sweep_baseline(base_factory,GRID,seeds=SWEEP_SEEDS) trials=[{'cfg':c,'result':evaluate(idea_factory(c),SEEDS)} for c in GRID] best=min(trials,key=lambda q:q['result']['mean']) sig=[run(best['cfg'],s,True,True)[1] for s in SEEDS] report=make_report('tabular','mlp_tiny',base,best['result'],{ 'idea_config':best['cfg'],'idea_sweep':trials, 'mechanism_signature':{'prediction':'training lowers symmetric nearest cross-color embedding distance','predicted_direction':'lower','observed_idea_cross_color_nearest_mean':float(np.mean(sig)),'observed_idea_cross_color_nearest_per_seed':sig,'confirmed':bool(np.isfinite(sig).all())} }) report['mechanism_signature']=report.pop('mechanism_signature') with open('bench_report.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__=='__main__': main()