Cyclic Lie-Bracket Residual Block / bench_cyclic.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()