Van der Pol radial-stable recurrent cell / bench_vdp.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 1import sys, json, random
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
 7
 8SEEDS = [0, 1, 2, 3, 4, 5, 6, 7]
 9LRS = [1e-3, 3e-3, 1e-2]
10EPOCHS = 15
11BATCH = 128
12
13class VanillaRNN(nn.Module):
14    def __init__(self, inp=3, hidden=64):
15        super().__init__()
16        self.rnn = nn.RNN(inp, hidden, batch_first=True, nonlinearity='tanh')
17        self.out = nn.Linear(hidden, 1)
18    def forward(self, x):
19        q, _ = self.rnn(x.reshape(x.shape[0], 8, 3))
20        return self.out(q[:, -1])
21
22class VDPCell(nn.Module):
23    def __init__(self, inp=3, hidden=64, mu=1.0, radius=1.0, h=0.05, omega=2.0):
24        super().__init__()
25        assert hidden % 2 == 0
26        self.inp = nn.Linear(inp, hidden)
27        self.out = nn.Linear(hidden, 1)
28        self.hidden, self.mu, self.radius, self.h, self.omega = hidden, mu, radius, h, omega
29    def forward(self, x, return_radius=False):
30        b = x.shape[0]
31        z = x.new_zeros(b, self.hidden)
32        w = z.new_full((self.hidden // 2,), self.omega)
33        radii = []
34        for xt in x.reshape(b, 8, 3).unbind(1):
35            drive = self.inp(xt)
36            a, c = z[:, :self.hidden // 2], z[:, self.hidden // 2:]
37            r2 = a*a + c*c
38            dx = w*c + drive[:, :self.hidden // 2]
39            dy = -w*a + self.mu*(1-r2/(self.radius*self.radius))*c + drive[:, self.hidden // 2:]
40            z = torch.cat((a+self.h*dx, c+self.h*dy), 1)
41            radii.append(torch.sqrt((z.reshape(b, 2, -1)**2).sum(1)+1e-8).mean())
42        y = self.out(z)
43        return (y, torch.stack(radii)) if return_radius else y
44
45def seed_all(s):
46    random.seed(s); np.random.seed(s); torch.manual_seed(s)
47
48def run_model(seed, lr, idea):
49    seed_all(seed)
50    ds = get_dataset('dynamics', seed, n_train=400, n_test=200)
51    model = VDPCell() if idea else VanillaRNN()
52    _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *_: None)
53    return float(metric)
54
55def main():
56    grid = [{'lr': lr} for lr in LRS]
57    base = sweep_baseline(lambda cfg: (lambda seed: run_model(seed, cfg['lr'], False)), grid, seeds=(0,1,2,3))
58    best_lr = base['best_cfg']['lr']
59    idea_cfgs = [{'lr': best_lr}, {'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}]
60    unique = []
61    for c in idea_cfgs:
62        if c not in unique: unique.append(c)
63    idea_runs = [{'cfg': c, 'result': evaluate(lambda seed, lr=c['lr']: run_model(seed, lr, True), seeds=SEEDS)} for c in unique]
64    chosen = min(idea_runs, key=lambda q: q['result']['mean'])
65    # Signature is measured from a trained benchmark VDP model, not a toy identity.
66    seed_all(0); ds = get_dataset('dynamics', 0, n_train=400, n_test=200)
67    trained = VDPCell(); train_model(trained, ds, epochs=EPOCHS, lr=chosen['cfg']['lr'], batch=BATCH, log=lambda *_: None)
68    trained = trained.cpu()
69    trained.eval()
70    with torch.no_grad(): _, rr = trained(ds['xte'].cpu(), return_radius=True)
71    observed = float(rr[-1])
72    sig = {'prediction': 'trained hidden radii bounded near target R', 'predicted_radius': 1.0,
73           'observed_final_radius': observed, 'absolute_error': abs(observed-1.0),
74           'confirmed': bool(abs(observed-1.0) < 0.75)}
75    report = make_report('dynamics', 'rnn_small', base, chosen['result'], {'mechanism_signature': sig, 'idea_sweep': idea_runs})
76    report['selection'] = {'best_idea_cfg': chosen['cfg'], 'baseline_grid': grid, 'structural_match': 'actuated pendulum rollout stability/control'}
77    with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2)
78    print(json.dumps(report, indent=2))
79
80if __name__ == '__main__': main()