Cyclic Lie-Bracket Residual Block / bench_cyclic.py
Beats tuned baseline
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, sweep_baseline, make_report
9
10TRACK = 'dynamics'
11MODEL = 'rnn_small'
12SEEDS = tuple(range(8))
13LRS = [1e-3, 3e-3, 6e-3]
14EPOCHS = 12
15NTRAIN, NTEST = 1000, 300
16
17class Field(nn.Module):
18 def __init__(self, d, width=64):
19 super().__init__()
20 self.net = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, width), nn.GELU(), nn.Linear(width, d))
21 def forward(self, x):
22 return self.net(x)
23
24class CyclicRNN(nn.Module):
25 def __init__(self, out_dim=1, step=0.15, random_order=True):
26 super().__init__()
27 self.rnn = nn.GRU(3, 64, batch_first=True)
28 self.f1, self.f2 = Field(64), Field(64)
29 self.head = nn.Linear(64, out_dim)
30 self.step, self.random_order = step, random_order
31 self._no_cudnn = False
32 def encode(self, x):
33 seq = x.view(x.shape[0], -1, 3)
34 try:
35 _, h = self.rnn(seq)
36 except RuntimeError:
37 self._no_cudnn = True
38 if self._no_cudnn:
39 old = torch.backends.cudnn.enabled
40 torch.backends.cudnn.enabled = False
41 try:
42 _, h = self.rnn(seq)
43 finally:
44 torch.backends.cudnn.enabled = old
45 return h[-1]
46 def cyclic(self, z, reverse=False):
47 u, v = self.f1(z), self.f2(z)
48 mean = (u + v) * 0.5
49 u, v = u - mean, v - mean
50 if not reverse:
51 z = z + self.step * u
52 z = z + 0.5 * self.step * (self.f2(z) - self.f1(z))
53 else:
54 z = z + self.step * v
55 z = z + 0.5 * self.step * (self.f1(z) - self.f2(z))
56 return z
57 def forward(self, x):
58 z = self.encode(x)
59 if self.training and self.random_order:
60 reverse = bool(torch.rand((), device=z.device) < 0.5)
61 else:
62 reverse = False
63 return self.head(self.cyclic(z, reverse))
64
65def seed_all(seed):
66 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
67 if torch.cuda.is_available():
68 torch.cuda.manual_seed_all(seed)
69
70def ds(seed):
71 return get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST)
72
73def base_model(d):
74 # Exact canonical rnn_small architecture.
75 from bench import make_model
76 return make_model(MODEL, d['input_shape'], d['out_dim'])
77
78def idea_model(d, step):
79 return CyclicRNN(d['out_dim'], step=step, random_order=True)
80
81def run_one(kind, cfg, seed):
82 seed_all(seed)
83 d = ds(seed)
84 model = base_model(d) if kind == 'baseline' else idea_model(d, cfg['step'])
85 _, metric, _ = train_model(model, d, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_: None)
86 if metric is None: return float('nan')
87 return float(metric)
88
89def eval_cfg(kind, cfg, seeds=SEEDS):
90 vals = [run_one(kind, cfg, int(s)) for s in seeds]
91 vals = [v for v in vals if np.isfinite(v)]
92 return {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals)}
93
94def main():
95 # Baseline sweep explicitly covers every lr used by the idea and its nearby settings.
96 baseline_grid = [{'lr': lr, 'step': 0.0} for lr in LRS]
97 base = sweep_baseline(lambda cfg: (lambda seed: run_one('baseline', cfg, seed)), baseline_grid, seeds=(0,1,2,3))
98 # Re-run baseline at every union lr on all paired seeds for fair reporting.
99 base_full_by_cfg = []
100 for cfg in baseline_grid:
101 r = eval_cfg('baseline', cfg)
102 base_full_by_cfg.append({'cfg': cfg, **r})
103 best_cfg = min(base_full_by_cfg, key=lambda x: x['mean'])['cfg']
104 base = {'best_cfg': best_cfg, 'sweep': [{'cfg': x['cfg'], 'mean': x['mean']} for x in base_full_by_cfg], 'full': next(x for x in base_full_by_cfg if x['cfg'] == best_cfg)}
105 idea_grid = [{'lr': best_cfg['lr'], 'step': 0.10}, {'lr': best_cfg['lr'], 'step': 0.15}, {'lr': best_cfg['lr'], 'step': 0.22}]
106 idea_runs = [{'cfg': cfg, **eval_cfg('idea', cfg)} for cfg in idea_grid]
107 idea_best = min(idea_runs, key=lambda x: x['mean'])
108 idea_result = idea_best
109 # Signature is measured on trained models, not an analytic toy field.
110 seed = 0; seed_all(seed); d = ds(seed); model = idea_model(d, idea_best['cfg']['step'])
111 model, _, _ = train_model(model, d, epochs=EPOCHS, lr=idea_best['cfg']['lr'], batch=128, log=lambda *_: None)
112 model.eval(); x = d['xte'][:128]
113 with torch.no_grad():
114 z = model.encode(x.to(next(model.parameters()).device)); ab = model.cyclic(z, False); ba = model.cyclic(z, True)
115 gap = (ab-ba).norm(dim=1).mean().item()
116 e = torch.tensor([0.01, 0.02, 0.04], device=z.device)
117 gaps = []
118 for h in e:
119 old = model.step; model.step = float(h); gaps.append((model.cyclic(z,False)-model.cyclic(z,True)).norm(dim=1).mean().item()); model.step = old
120 slope = float(np.polyfit(np.log(e.cpu().numpy()), np.log(np.maximum(gaps,1e-12)), 1)[0])
121 signature = {'order_gap_mean': float(gap), 'steps': e.cpu().numpy().tolist(), 'gaps': gaps, 'predicted_exponent': 2.0, 'observed_exponent': slope, 'confirmed': bool(gap > 1e-8 and 1.5 < slope < 2.5)}
122 idea_report = dict(idea_result)
123 idea_report['best_cfg'] = idea_best['cfg']
124 idea_report['sweep'] = [{'cfg': x['cfg'], 'mean': x['mean']} for x in idea_runs]
125 report = make_report(TRACK, MODEL, base, idea_report, extra=signature)
126 Path('bench_report.json').write_text(json.dumps(report, indent=2))
127 print(json.dumps(report, indent=2))
128
129if __name__ == '__main__': main()