Mpemba Mode-Filtered Training / mpemba_bench.py
Mechanism confirmed, baseline not beaten
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()