Fisher-floor-corrected DSM / bench_exp.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
 8
 9TRACK='multitoken_diffusion'
10MODEL='mlp_tiny'
11DIM=8
12BETA=8.0
13
14def floor_estimate(noisy_t, bank, t):
15    # Gaussian VP corruption: alpha=exp(-beta*t/2), sigma^2=1-alpha^2.
16    alpha=torch.exp(-BETA*t/2).unsqueeze(1)
17    sigma2=(1-alpha*alpha).clamp_min(1e-5)
18    logits=-((noisy_t[:,None,:]-alpha[:,None,:]*bank[None,:,:])**2).sum(-1)/(2*sigma2)
19    pi=torch.softmax(logits, dim=1)
20    mu=(pi[:,:,None]*bank[None,:,:]).sum(1)
21    var=(pi*((bank[None,:,:]-mu[:,None,:])**2).sum(-1)).sum(1)
22    return (alpha[:,0]**2/(sigma2[:,0]**2))*var
23
24def train(seed, cfg, corrected):
25    torch.manual_seed(seed); np.random.seed(seed)
26    ds=get_dataset(TRACK, seed, n_train=400, n_test=200)
27    model=make_model(MODEL, ds['input_shape'], DIM)
28    device='cuda' if torch.cuda.is_available() else 'cpu'
29    try:
30        model=model.to(device)
31        opt=torch.optim.Adam(model.parameters(), lr=cfg['lr'])
32        x=ds['xtr'].to(device); clean=ds['ytr'].reshape(-1,DIM).to(device)
33        bank=clean.detach().clone()
34        hist=[]
35        for ep in range(cfg['epochs']):
36            model.train(); perm=torch.randperm(len(x),device=device); total=0.
37            for i in range(0,len(x),128):
38                ix=perm[i:i+128]; inp=x[ix]; y=clean[ix]
39                noisy=inp[:,:DIM]; t=inp[:,DIM]
40                alpha=torch.exp(-BETA*t/2).unsqueeze(1)
41                sigma2=(1-alpha*alpha).clamp_min(1e-5)
42                target=(alpha*y-noisy)/sigma2
43                raw=((model(inp)-target)**2).sum(1)
44                if corrected:
45                    floor=floor_estimate(noisy,bank,t).detach()
46                    loss=(raw-floor).mean()
47                else: loss=raw.mean()
48                opt.zero_grad(); loss.backward(); opt.step(); total += float(loss.detach())*len(ix)
49            hist.append(total/len(x))
50        model.eval()
51        with torch.no_grad():
52            # Posterior-mean denoising readout from predicted score; standard task MSE.
53            inp=ds['xte'].to(device); noisy=inp[:,:DIM]; t=inp[:,DIM]
54            alpha=torch.exp(-BETA*t/2).unsqueeze(1); sigma2=(1-alpha*alpha).clamp_min(1e-5)
55            score=model(inp); pred=(noisy+sigma2*score)/alpha.clamp_min(1e-4)
56            metric=float(((pred-ds['yte'].reshape(-1,DIM).to(device))**2).mean())
57            # behavior signature: gradient equality and output agreement versus a fresh batch
58            raw=((score-((alpha*ds['yte'].reshape(-1,DIM).to(device)-noisy)/sigma2))**2).sum(1)
59        return metric, {'final_train_loss':hist[-1], 'mean_raw_dsm':float(raw.mean()), 'pred_std':float(pred.std())}
60    except RuntimeError:
61        # CPU retry is explicit per bench requirement.
62        torch.cuda.empty_cache() if torch.cuda.is_available() else None
63        device='cpu'; model=make_model(MODEL, ds['input_shape'], DIM).to(device)
64        opt=torch.optim.Adam(model.parameters(),lr=cfg['lr']); x=ds['xtr']; clean=ds['ytr'].reshape(-1,DIM); bank=clean.detach().clone()
65        for _ in range(cfg['epochs']):
66            for i in range(0,len(x),128):
67                inp=x[i:i+128]; y=clean[i:i+128]; noisy=inp[:,:DIM]; t=inp[:,DIM]
68                a=torch.exp(-BETA*t/2).unsqueeze(1); s=(1-a*a).clamp_min(1e-5); target=(a*y-noisy)/s
69                raw=((model(inp)-target)**2).sum(1); loss=(raw-floor_estimate(noisy,bank,t).detach()).mean() if corrected else raw.mean()
70                opt.zero_grad(); loss.backward(); opt.step()
71        with torch.no_grad():
72            inp=ds['xte']; noisy=inp[:,:DIM]; t=inp[:,DIM]; a=torch.exp(-BETA*t/2).unsqueeze(1); s=(1-a*a).clamp_min(1e-5)
73            pred=(noisy+s*model(inp))/a.clamp_min(1e-4); metric=float(((pred-ds['yte'].reshape(-1,DIM))**2).mean())
74        return metric, {'fallback':'cpu'}
75
76def factory(cfg, corrected):
77    return lambda seed: train(seed,cfg,corrected)[0]
78
79def main():
80    # Shared union: every lr tested by the idea is also swept for baseline.
81    grid=[{'lr':1e-3,'epochs':15},{'lr':3e-3,'epochs':15},{'lr':6e-3,'epochs':15}]
82    small=tuple(range(8))
83    base=sweep_baseline(lambda c:factory(c,False),grid,seeds=small)
84    bcfg=base['best_cfg']
85    base_full=evaluate(factory(bcfg,False),seeds=small)
86    idea_runs=[(c,evaluate(factory(c,True),seeds=small)) for c in grid]
87    icfg, idea=min(idea_runs,key=lambda z:z[1]['mean'])
88    # NN-scale mechanism signature: measure loss invariance and observed paired outputs.
89    probe=[]
90    for s in small:
91        a,_=train(s,icfg,False); b,_=train(s,icfg,True); probe.append({'seed':s,'baseline_metric':a,'idea_metric':b,'delta':b-a})
92    signature={'prediction':'detached floor subtraction leaves parameter gradients unchanged under identical weighting','predicted_vs_observed':{'predicted_gradient_change':0.0,'observed_metric_delta_mean':float(np.mean([p['delta'] for p in probe])),'observed_abs_delta_max':float(np.max(np.abs([p['delta'] for p in probe])))},'trained_model_observations':probe,'confirmed':False}
93    report=make_report(TRACK,MODEL,{'best_cfg':bcfg,'sweep':base['sweep'],'full':base_full},idea,signature)
94    report['idea_sweep']=[{'cfg':c,'result':r} for c,r in idea_runs]
95    report['custom_track']={'name':TRACK,'file':'bench_exp.py','domain':'diffusion-sampling'}
96    Path('bench_report.json').write_text(json.dumps(report,indent=2))
97    print(json.dumps(report,indent=2))
98
99if __name__=='__main__': main()