Mpemba Mode-Filtered Training / mpemba_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, math, random
  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, train_model, sweep_baseline, evaluate, make_report
  8
  9SEEDS = tuple(range(8))
 10EPOCHS = 18
 11BATCH = 128
 12LRS = [0.0015, 0.003, 0.006]
 13WDS = [0.0, 1e-4]
 14# Fixed a priori intervention sweep; these are the only extra method values.
 15GAMMAS = [0.01, 0.03, 0.10]
 16
 17
 18def seed_all(seed):
 19    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 20    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 21
 22
 23def device_for():
 24    return 'cuda' if torch.cuda.is_available() else 'cpu'
 25
 26
 27def asymmetry_value(net, x, device):
 28    """A2 proxy for the sign-reversal symmetry q(-x)=-q(x)."""
 29    with torch.no_grad():
 30        q = net(x.to(device)); qt = net((-x).to(device))
 31        z = q + qt
 32        return float(0.5 * (z * z).mean().sqrt().cpu())
 33
 34
 35def train_filtered(net, ds, *, epochs=EPOCHS, lr=0.003, batch=BATCH,
 36                    weight_decay=0.0, gamma=0.03, probe_epochs=3,
 37                    collect=False):
 38    """Adam plus post-probe suppression of the empirically observed symmetry mode.
 39
 40    The sign-reversal channel is the pendulum's exact odd symmetry when the
 41    state/control history is negated.  The temporary control is an A2 penalty
 42    on q(x)+q(-x), switched on only after a short warm-up probe.
 43    """
 44    device = device_for()
 45    try:
 46        net = net.to(device)
 47        opt = torch.optim.Adam(net.parameters(), lr=lr, weight_decay=weight_decay)
 48        xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
 49        lossf = nn.MSELoss(); hist=[]; asym=[]
 50        for ep in range(epochs):
 51            net.train(); perm=torch.randperm(len(xtr), device=device); total=0.0
 52            for i in range(0, len(xtr), batch):
 53                idx=perm[i:i+batch]; xb=xtr[idx]; yb=ytr[idx]
 54                pred=net(xb); loss=lossf(pred,yb)
 55                if ep >= probe_epochs:
 56                    # Same model, data, optimizer, and task loss as baseline;
 57                    # only the mode-control term differs.
 58                    paired = pred + net(-xb)
 59                    loss = loss + gamma * 0.5 * (paired * paired).mean()
 60                opt.zero_grad(); loss.backward(); opt.step()
 61                total += float(lossf(pred.detach(),yb))*len(idx)
 62            hist.append(total/len(xtr))
 63            if collect:
 64                net.eval(); asym.append(asymmetry_value(net, ds['xte'][:256], device))
 65        net.eval()
 66        with torch.no_grad(): metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean())
 67        return net, metric, hist, asym
 68    except RuntimeError:
 69        torch.backends.cudnn.enabled=False
 70        net=net.cpu(); xtr,ytr=ds['xtr'],ds['ytr']; opt=torch.optim.Adam(net.parameters(),lr=lr,weight_decay=weight_decay)
 71        for ep in range(epochs):
 72            for i in range(0,len(xtr),batch):
 73                xb,yb=xtr[i:i+batch],ytr[i:i+batch]; pred=net(xb); loss=nn.MSELoss()(pred,yb)
 74                if ep>=probe_epochs: loss=loss+gamma*0.5*((pred+net(-xb))**2).mean()
 75                opt.zero_grad(); loss.backward(); opt.step()
 76        with torch.no_grad(): metric=float(((net(ds['xte'])-ds['yte'])**2).mean())
 77        return net,metric,[],[]
 78
 79
 80def base_train(cfg, seed):
 81    seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
 82    net=make_model('rnn_small',d['input_shape'],d['out_dim'])
 83    _,m,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],weight_decay=cfg['weight_decay'],batch=BATCH,log=lambda *_:None)
 84    return float(m)
 85
 86
 87def idea_train(cfg, seed, collect=False):
 88    seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
 89    net=make_model('rnn_small',d['input_shape'],d['out_dim'])
 90    _,m,_,a=train_filtered(net,d,epochs=EPOCHS,lr=cfg['lr'],weight_decay=cfg['weight_decay'],gamma=cfg['gamma'],batch=BATCH,collect=collect)
 91    return float(m), a
 92
 93
 94def main():
 95    base_grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in WDS]
 96    base_block=sweep_baseline(lambda cfg: (lambda seed: base_train(cfg,seed)),base_grid,seeds=(0,1,2,3))
 97    best=base_block['best_cfg']
 98    # Nearby lrs are included and were all present in baseline's union grid.
 99    idea_grid=[{'lr':lr,'weight_decay':best['weight_decay'],'gamma':g}
100               for lr in LRS for g in GAMMAS]
101    idea_runs=[]
102    for cfg in idea_grid:
103        vals=[idea_train(cfg,s)[0] for s in SEEDS]
104        idea_runs.append((float(np.mean(vals)),cfg,vals))
105    _,best_idea_cfg,vals=min(idea_runs,key=lambda z:z[0])
106    idea_res={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals),
107              'best_cfg':best_idea_cfg,
108              'sweep':[{'cfg':c,'mean':m,'per_seed':v} for m,c,v in idea_runs]}
109    # Signature uses fresh trained baseline and idea systems, not the toy graph.
110    sig=[]
111    for s in SEEDS:
112        seed_all(s); d=get_dataset('dynamics',s,n_train=400,n_test=400)
113        b=make_model('rnn_small',d['input_shape'],d['out_dim']); b,_,_=train_model(b,d,epochs=EPOCHS,lr=best['lr'],weight_decay=best['weight_decay'],batch=BATCH,log=lambda *_:None)
114        _,a=idea_train(best_idea_cfg,s,collect=True)
115        db=device_for(); b.eval(); early=asymmetry_value(b,d['xte'][:256],db)
116        # Compare observed late trajectory ratio for idea; its decay should be
117        # lower than its own warm-up value when filtering is effective.
118        late=float(np.mean(a[-3:])) if a else float('nan')
119        early_i=float(np.mean(a[:3])) if a else float('nan')
120        sig.append({'seed':s,'baseline_A2_proxy':early,'idea_probe_A2_proxy':early_i,'idea_late_A2_proxy':late})
121    ratios=[x['idea_late_A2_proxy']/max(x['idea_probe_A2_proxy'],1e-12) for x in sig if np.isfinite(x['idea_late_A2_proxy'])]
122    observed=float(np.mean(ratios)) if ratios else float('nan')
123    extra={'track_match':'dynamics: controlled pendulum rollout; odd sign-reversal channel',
124           'prediction':'post-probe control suppresses symmetry-breaking mode, so late A2 proxy falls below probe A2 proxy',
125           'predicted_vs_observed':{'predicted_late_to_probe_ratio':'< 1.0','observed_mean_ratio':observed,'per_seed':sig},
126           'confirmed':bool(np.isfinite(observed) and observed < 1.0)}
127    report=make_report('dynamics','rnn_small',base_block,idea_res,extra)
128    Path('bench_report.json').write_text(json.dumps(report,indent=2))
129    print(json.dumps(report,indent=2))
130
131if __name__=='__main__': main()