Geometrically Attracting Random Recurrent Layer / bench_run.py
Failed on benchmark
1import os, sys, json, math
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
7from bench.protocol import evaluate
8
9SEED0 = 2803
10EPOCHS = 12
11NTR, NTE = 1200, 400
12LRS = [1e-3, 3e-3, 1e-2]
13TARGET_RHO = 0.95
14LAMBDA = 0.05
15
16class RandomAttractingRegressor(nn.Module):
17 def __init__(self, hidden=64, candidates=2, target_rho=0.95):
18 super().__init__()
19 self.hidden, self.k, self.target_rho = hidden, candidates, target_rho
20 self.W = nn.Parameter(torch.empty(candidates, hidden, hidden))
21 self.U = nn.Parameter(torch.empty(candidates, hidden, 3))
22 self.b = nn.Parameter(torch.zeros(candidates, hidden))
23 self.gate = nn.Linear(3, candidates)
24 self.head = nn.Linear(hidden, 1)
25 for i in range(candidates):
26 nn.init.orthogonal_(self.W[i])
27 with torch.no_grad():
28 self.W[0].mul_(0.82); self.W[1].mul_(1.14)
29 nn.init.xavier_uniform_(self.U)
30 nn.init.zeros_(self.gate.weight)
31 nn.init.constant_(self.gate.bias, 0.0)
32
33 def gains(self):
34 return torch.linalg.matrix_norm(self.W, ord=2, dim=(-2, -1))
35
36 def forward(self, x, return_aux=False, h0=None):
37 # x is [batch, 24], eight (theta, omega, control) observations.
38 seq = x.view(x.shape[0], -1, 3)
39 h = x.new_zeros(x.shape[0], self.hidden) if h0 is None else h0
40 states, probs = [], []
41 for t in range(seq.shape[1]):
42 xt = seq[:, t]
43 p = torch.softmax(self.gate(xt), dim=-1)
44 cand = torch.tanh(torch.einsum('kij,bj->bki', self.W, h) +
45 torch.einsum('kij,bj->bki', self.U, xt) + self.b)
46 h = (p.unsqueeze(-1) * cand).sum(1)
47 states.append(h); probs.append(p)
48 out = self.head(h)
49 if return_aux:
50 return out, torch.stack(states, 1), torch.stack(probs, 1)
51 return out
52
53 def contraction_penalty(self, probs):
54 eg = (probs * self.gains()).sum(-1)
55 excess = torch.log(eg + 1e-8) - math.log(self.target_rho)
56 return torch.relu(excess).square().mean()
57
58def seed_all(seed):
59 np.random.seed(seed); torch.manual_seed(seed)
60 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
61
62def data(seed):
63 return get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
64
65def train_random(cfg, seed, lam, keep=False):
66 seed_all(SEED0 + seed * 101 + int(cfg['lr'] * 1e6))
67 ds = data(seed)
68 model = RandomAttractingRegressor(target_rho=TARGET_RHO)
69 devices = ['cuda', 'cpu'] if torch.cuda.is_available() else ['cpu']
70 for dev in devices:
71 try:
72 net = model.to(dev)
73 opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'])
74 xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
75 for ep in range(EPOCHS):
76 net.train(); perm = torch.randperm(len(xtr), device=dev)
77 for j in range(0, len(xtr), 128):
78 ix = perm[j:j+128]
79 pred, _, p = net(xtr[ix], return_aux=True)
80 loss = ((pred-ytr[ix])**2).mean() + lam * net.contraction_penalty(p)
81 opt.zero_grad(); loss.backward(); opt.step()
82 net.eval()
83 with torch.no_grad():
84 pred = net(ds['xte'].to(dev))
85 metric = float(((pred-ds['yte'].to(dev)) ** 2).mean())
86 if keep:
87 torch.save(net.state_dict(), 'idea_model.pt' if lam else 'baseline_model.pt')
88 return metric, net, ds, dev
89 except RuntimeError:
90 if dev == 'cuda':
91 continue
92 return float('nan'), None, ds, 'cpu'
93
94def baseline_metric(cfg, seed):
95 return train_random(cfg, seed, 0.0)[0]
96
97def idea_train(cfg, seed, keep=False):
98 return train_random(cfg, seed, LAMBDA, keep)
99
100def idea_metric(cfg, seed):
101 return idea_train(cfg, seed)[0]
102
103def signature(cfg, seed=0):
104 metric, net, ds, dev = idea_train(cfg, seed, keep=True)
105 if net is None: return {'confirmed': False, 'error': 'training failed'}
106 x = ds['xte'][:64].to(dev)
107 with torch.no_grad():
108 _, states, probs = net(x, return_aux=True)
109 # Same inputs, two nearby initial states; re-run explicitly for observed contraction.
110 h0 = torch.zeros(x.shape[0], net.hidden, device=dev); h1 = h0.clone(); h1[:,0] = 1.0
111 _, s0, _ = net(x, return_aux=True, h0=h0)
112 _, s1, _ = net(x, return_aux=True, h0=h1)
113 d = (s0-s1).norm(dim=-1).mean(0).cpu().numpy() + 1e-12
114 slope = float(np.polyfit(np.arange(len(d))[-4:], np.log(d)[-4:], 1)[0])
115 gains = net.gains().cpu().numpy()
116 pp = probs.mean((0,1)).cpu().numpy()
117 predicted = float(np.log(np.sum(pp*gains)))
118 return {'predicted_log_expected_gain': predicted, 'observed_log_distance_slope': slope,
119 'relative_slope_error': abs(slope-predicted)/(abs(predicted)+1e-8),
120 'mean_route_probabilities': pp.tolist(), 'candidate_spectral_gains': gains.tolist(),
121 'test_mse': metric, 'confirmed': bool(abs(slope-predicted)/(abs(predicted)+1e-8) < 0.35)}
122
123def main():
124 grid = [{'lr': v} for v in LRS]
125 base = sweep_baseline(lambda c: lambda s: baseline_metric(c, s), grid)
126 idea_runs = []
127 for cfg in grid:
128 r = evaluate(lambda s, c=cfg: idea_metric(c, s))
129 idea_runs.append({'cfg': cfg, 'result': r})
130 best = min(idea_runs, key=lambda z: z['result']['mean'])
131 idea = best['result']
132 sig = signature(best['cfg'], 0)
133 report = make_report('dynamics', 'rnn_small', base, idea,
134 {'predicted_vs_observed': sig, 'idea_sweep': idea_runs,
135 'track_justification': 'Dynamics is the built-in structural match for recurrent stability/control.'})
136 report['protocol'] = {'epochs': EPOCHS, 'n_train': NTR, 'n_test': NTE,
137 'lr_union': LRS, 'idea_lambda': LAMBDA, 'target_rho': TARGET_RHO}
138 with open('bench_report.json','w') as f: json.dump(report, f, indent=2)
139 print(json.dumps(report, indent=2))
140
141if __name__ == '__main__': main()