Mean-Square Proximal Relaxation Optimizer / stage2_bench.py
Failed on benchmark
1import json, random, sys
2import numpy as np
3import torch
4from torch import nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report, count_params
8
9SEEDS = tuple(range(8))
10SWEEP_SEEDS = tuple(range(4))
11EPOCHS = 15
12BATCH = 128
13LR_GRID = [1e-3, 3e-3, 6e-3]
14ALPHA_GRID = [0.25, 0.5, 0.75]
15
16
17def seed_all(seed):
18 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
19 if torch.cuda.is_available():
20 try: torch.cuda.manual_seed_all(seed)
21 except Exception: pass
22
23
24def device():
25 return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
26
27
28def proximal_train(model, ds, lr, alpha, epochs=EPOCHS, batch=BATCH, return_stats=False):
29 """Blockwise stochastic proximal response T=w-lr*g, then w<-w+alpha(T-w).
30 Blocks are parameter tensors; gradients are computed jointly but applied blockwise,
31 making the intervention the optimizer update rather than the network architecture.
32 """
33 try:
34 dev = device(); model = model.to(dev)
35 xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
36 except Exception:
37 dev = torch.device('cpu'); model = model.to(dev)
38 xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
39 loss_fn = nn.MSELoss()
40 params = [p for p in model.parameters() if p.requires_grad]
41 rng = np.random.default_rng(12345)
42 update_norms = []; grad_snapshots = []
43 n = len(xtr)
44 model.train()
45 for ep in range(epochs):
46 order = rng.permutation(n)
47 for start in range(0, n, batch):
48 ix = torch.as_tensor(order[start:start+batch], device=dev)
49 model.zero_grad(set_to_none=True)
50 pred = model(xtr[ix])
51 loss = loss_fn(pred, ytr[ix])
52 loss.backward()
53 norms = []
54 with torch.no_grad():
55 for p in params: # each tensor is one proximal block
56 if p.grad is None: continue
57 # T_hat = p - lr*g; relaxed response is p - alpha*lr*g.
58 step = alpha * lr * p.grad
59 p.sub_(step)
60 norms.append(float(step.norm().detach().cpu()))
61 if norms:
62 update_norms.append(float(np.sqrt(np.sum(np.square(norms)))))
63 grad_snapshots.append(float(np.mean(np.square(norms))))
64 model.eval()
65 with torch.no_grad():
66 metric = float(loss_fn(model(ds['xte'].to(dev)), ds['yte'].to(dev)).cpu())
67 stats = {'update_rms': float(np.sqrt(np.mean(np.square(update_norms)))) if update_norms else 0.0,
68 'update_std': float(np.std(update_norms)) if update_norms else 0.0,
69 'n_updates': len(update_norms)}
70 return model, metric, stats
71
72
73def baseline_fn(cfg):
74 def run(seed):
75 seed_all(seed)
76 ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
77 model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
78 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
79 return float(metric)
80 return run
81
82
83def idea_fn(cfg, retain=None):
84 def run(seed):
85 seed_all(seed)
86 ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
87 model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
88 net, metric, stats = proximal_train(model, ds, cfg['lr'], cfg['alpha'])
89 if retain is not None: retain[seed] = (net, ds, stats)
90 return float(metric)
91 return run
92
93
94def main():
95 # Baseline includes every lr used by the idea; alpha is the idea-only relaxation knob.
96 base = sweep_baseline(baseline_fn, [{'lr': x} for x in LR_GRID], seeds=SWEEP_SEEDS)
97 idea_grid = [{'lr': lr, 'alpha': a} for lr, a in zip(LR_GRID, ALPHA_GRID)]
98 idea_sweep = []
99 for cfg in idea_grid:
100 vals = [idea_fn(cfg)(s) for s in SWEEP_SEEDS]
101 idea_sweep.append({'cfg': cfg, 'mean': float(np.mean(vals)), 'per_seed': vals})
102 best_cfg = min(idea_sweep, key=lambda z: z['mean'])['cfg']
103
104 retained = {}
105 idea_vals = [idea_fn(best_cfg, retained)(s) for s in SEEDS]
106 idea_res = {'mean': float(np.mean(idea_vals)), 'std': float(np.std(idea_vals)),
107 'per_seed': idea_vals, 'n': 8, 'chosen_cfg': best_cfg,
108 'sweep': idea_sweep}
109
110 base_vals = []
111 for s in SEEDS:
112 seed_all(s); ds = get_dataset('dynamics', s, n_train=4000, n_test=1000)
113 m = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
114 _, metric, _ = train_model(m, ds, epochs=EPOCHS, lr=base['best_cfg']['lr'], batch=BATCH, log=lambda *_: None)
115 base_vals.append(float(metric))
116 base['full'] = {'mean': float(np.mean(base_vals)), 'std': float(np.std(base_vals)),
117 'per_seed': base_vals, 'n': 8}
118
119 # Signature is measured on trained benchmark models: response/update noise and damping.
120 rms = [v[2]['update_rms'] for v in retained.values()]
121 std = [v[2]['update_std'] for v in retained.values()]
122 alpha = best_cfg['alpha']
123 # For the relaxed recursion, observed update magnitude should scale approximately alpha.
124 sig = {'prediction': 'relaxation damps stochastic block-response updates approximately linearly in alpha',
125 'predicted_alpha': alpha, 'observed_update_rms_mean': float(np.mean(rms)),
126 'observed_update_std_mean': float(np.mean(std)),
127 'predicted_vs_observed_relation': 'measured from trained rnn_small updates; alpha scaling was not quantitatively tested',
128 'confirmed': False}
129 report = make_report('dynamics', 'rnn_small', base, idea_res, {'mechanism_signature': sig})
130 report['mechanism_signature'] = sig
131 report['custom_track'] = None
132 with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2)
133 print(json.dumps(report, indent=2))
134
135if __name__ == '__main__': main()