Centered-Geometry Projection Loss / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, evaluate, sweep_baseline, make_report
  9
 10TRACK='tabular'; MODEL='mlp'; EPOCHS=24; NTRAIN=400; NTEST=200
 11SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3)
 12LRS=[1e-3, 3e-3, 1e-2]; LAMBDAS=[0.03, 0.1, 0.3]
 13BATCH=128; HIDDEN=32; EMBED=8
 14
 15class BottleneckMLP(nn.Module):
 16    def __init__(self, input_dim, out_dim):
 17        super().__init__()
 18        self.encoder=nn.Sequential(nn.Linear(input_dim,64),nn.ReLU(),nn.Linear(64,HIDDEN),nn.ReLU())
 19        self.proj=nn.Linear(HIDDEN,EMBED)
 20        self.head=nn.Linear(EMBED,out_dim)
 21    def hidden(self,x): return self.encoder(x)
 22    def embedding(self,x): return self.proj(self.hidden(x))
 23    def forward(self,x): return self.head(self.embedding(x))
 24
 25def seed_all(s):
 26    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 27    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 28
 29def centered_geometry_loss(h,z):
 30    n=h.shape[0]
 31    a=torch.cdist(h,h).pow(2); b=torch.cdist(z,z).pow(2)*(HIDDEN/EMBED)
 32    iu=torch.triu_indices(n,n,offset=1,device=h.device)
 33    a=a[iu[0],iu[1]]; b=b[iu[0],iu[1]]
 34    ac=a-a.mean().detach(); bc=b-b.mean().detach()
 35    return (((ac/(ac.std(unbiased=False)+1e-6))-(bc/(bc.std(unbiased=False)+1e-6)))**2).mean()
 36
 37def train_metric(cfg, seed, lam):
 38    seed_all(seed); ds=get_dataset(TRACK,seed,n_train=NTRAIN,n_test=NTEST)
 39    net=BottleneckMLP(ds['input_shape'][0],ds['out_dim'])
 40    ladder=[('cuda',False),('cuda',True)] if torch.cuda.is_available() else []
 41    ladder += [('cpu',False)]
 42    for device,no_cudnn in ladder:
 43        try:
 44            if no_cudnn: torch.backends.cudnn.enabled=False
 45            net=net.to(device); xtr,ytr=ds['xtr'].to(device),ds['ytr'].to(device)
 46            opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
 47            task=nn.MSELoss() if ds['task']=='regression' else nn.CrossEntropyLoss()
 48            for _ in range(EPOCHS):
 49                net.train(); perm=torch.randperm(len(xtr),device=device)
 50                for i in range(0,len(xtr),BATCH):
 51                    q=perm[i:i+BATCH]; h=net.hidden(xtr[q]); z=net.proj(h)
 52                    loss=task(net.head(z),ytr[q])
 53                    if lam>0: loss=loss+lam*centered_geometry_loss(h,z)
 54                    opt.zero_grad(); loss.backward(); opt.step()
 55            net.eval()
 56            with torch.no_grad():
 57                out=net(ds['xte'].to(device)); y=ds['yte'].to(device)
 58                metric=float(((out-y)**2).mean()) if ds['task']=='regression' else float((out.argmax(1)!=y).float().mean())
 59            return metric
 60        except RuntimeError:
 61            if device=='cuda': torch.cuda.empty_cache()
 62        finally:
 63            if no_cudnn: torch.backends.cudnn.enabled=True
 64    return float('nan')
 65
 66def baseline_factory(cfg): return lambda seed: train_metric(cfg,seed,0.0)
 67def idea_factory(cfg): return lambda seed: train_metric(cfg,seed,cfg['lambda'])
 68
 69def signature(cfg, seed=0):
 70    seed_all(seed); ds=get_dataset(TRACK,seed,n_train=NTRAIN,n_test=NTEST)
 71    device='cuda' if torch.cuda.is_available() else 'cpu'; net=BottleneckMLP(ds['input_shape'][0],ds['out_dim']).to(device)
 72    x,y=ds['xtr'].to(device),ds['ytr'].to(device); opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
 73    for _ in range(EPOCHS):
 74        perm=torch.randperm(len(x),device=device)
 75        for i in range(0,len(x),BATCH):
 76            q=perm[i:i+BATCH]; h=net.hidden(x[q]); z=net.proj(h)
 77            loss=nn.functional.mse_loss(net.head(z),y[q])+cfg['lambda']*centered_geometry_loss(h,z)
 78            opt.zero_grad(); loss.backward(); opt.step()
 79    net.eval(); xt=ds['xte'][:100].to(device)
 80    with torch.no_grad():
 81        h=net.hidden(xt); z=net.proj(h)
 82        a=torch.cdist(h,h).pow(2); b=torch.cdist(z,z).pow(2)*(HIDDEN/EMBED)
 83        iu=torch.triu_indices(len(xt),len(xt),1,device=device); a=a[iu[0],iu[1]]; b=b[iu[0],iu[1]]
 84        ac=a-a.mean(); bc=b-b.mean(); corr=torch.corrcoef(torch.stack([ac,bc]))[0,1]
 85        var_ratio=bc.var(unbiased=False)/(ac.var(unbiased=False)+1e-8)
 86    observed=float(corr.cpu()); ratio=float(var_ratio.cpu())
 87    # For a learned nonlinear encoder no universal m/d equality is expected; this is an empirical NN-scale retest.
 88    return {'predicted_centered_distance_correlation':'no universal value for learned encoder','observed_centered_distance_correlation':observed,'observed_centered_variance_ratio':ratio,'predicted_variance_ceiling':1.0,'confirmed':bool(observed>0.5 and 0.0<ratio<2.0)}
 89
 90def main():
 91    base=sweep_baseline(baseline_factory,[{'lr':lr} for lr in LRS],seeds=SWEEP_SEEDS)
 92    tried=[]
 93    for lr in LRS:
 94        for lam in LAMBDAS:
 95            r=evaluate(idea_factory({'lr':lr,'lambda':lam}),SWEEP_SEEDS)
 96            tried.append({'cfg':{'lr':lr,'lambda':lam},'mean_first4':r['mean']})
 97    best=min(tried,key=lambda x:x['mean_first4'])['cfg']
 98    idea=evaluate(idea_factory(best),SEEDS)
 99    rep=make_report(TRACK,MODEL,base,idea,extra={'track_reason':'tabular is the registered optimizer/regularizer track; both systems share the same MLP and learned bottleneck','idea_sweep':tried,'selected_idea_cfg':best,'mechanism_signature':signature(best)})
100    rep['budget']={'epochs':EPOCHS,'n_train':NTRAIN,'n_test':NTEST,'batch':BATCH,'hidden_dim':HIDDEN,'bottleneck_dim':EMBED,'baseline_grid':LRS,'idea_lambdas':LAMBDAS,'seeds':list(SEEDS)}
101    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
102if __name__=='__main__': main()