Event-triggered phase desynchronisation for recurrent hidden states / stage2_bench.py
Unverified
1import json, sys
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8import bench
9
10SEEDS = tuple(range(8))
11GRID = [
12 {'lr': 0.001, 'weight_decay': 0.0},
13 {'lr': 0.003, 'weight_decay': 0.0},
14 {'lr': 0.006, 'weight_decay': 0.0},
15]
16
17
18def control(z, k):
19 r = z.mean(1, keepdim=True)
20 return (2.0 * k / z.shape[1]) * torch.imag(z * torch.conj(r))
21
22
23def math_check(seed=123):
24 rng = np.random.default_rng(seed)
25 th = torch.tensor(rng.normal(0, .22, 24), dtype=torch.float64, requires_grad=True)
26 z = torch.exp(1j * th)
27 V = torch.abs(z.mean()) ** 2
28 grad = torch.autograd.grad(V, th)[0]
29 dt = 1e-7
30 th2 = th.detach() + dt * (-grad.detach())
31 V2 = torch.abs(torch.exp(1j * th2).mean()) ** 2
32 observed = float((V2 - V.detach()) / dt)
33 predicted = float(-(grad * grad).sum())
34 ratio = observed / predicted
35 return {'Vdot_observed': observed, 'Vdot_predicted': predicted,
36 'ratio': ratio, 'passed': abs(ratio - 1) < 1e-4}
37
38
39class PhaseGRU(nn.Module):
40 """The same GRUCell architecture in both arms; event rotation is the only change."""
41 def __init__(self, hidden=64, mode='baseline', k=1.0, delta=0.05):
42 super().__init__()
43 if hidden % 2:
44 raise ValueError('hidden must be even')
45 self.cell = nn.GRUCell(3, hidden)
46 self.head = nn.Linear(hidden, 1)
47 self.mode, self.k, self.delta = mode, k, delta
48 self.last_signature = {}
49
50 def forward(self, x, collect=False):
51 b = x.shape[0]
52 h = x.new_zeros(b, self.cell.hidden_size)
53 held = x.new_zeros(b, self.cell.hidden_size // 2)
54 vs, exact_norms, held_norms, events, errors = [], [], [], [], []
55 for t in range(x.shape[1] // 3):
56 h = self.cell(x[:, 3*t:3*t+3], h)
57 p = h.view(b, -1, 2)
58 norm = torch.sqrt((p*p).sum(-1) + 1e-8)
59 z = torch.complex(p[..., 0], p[..., 1]) / norm
60 u = control(z, self.k)
61 if self.mode == 'continuous':
62 applied = u
63 ev = torch.ones(b, device=x.device, dtype=torch.bool)
64 elif self.mode == 'event':
65 ev = torch.linalg.vector_norm(u - held, dim=1) >= self.delta
66 applied = torch.where(ev[:, None], u, held)
67 held = applied
68 else:
69 applied = torch.zeros_like(u)
70 ev = torch.zeros(b, device=x.device, dtype=torch.bool)
71 angle = applied.unsqueeze(-1)
72 rot = torch.stack((p[..., 0]*torch.cos(angle[..., 0]) - p[..., 1]*torch.sin(angle[..., 0]),
73 p[..., 0]*torch.sin(angle[..., 0]) + p[..., 1]*torch.cos(angle[..., 0])), -1)
74 h = rot.reshape(b, -1)
75 if collect:
76 r = z.mean(1)
77 vs.append((r.abs()**2).detach())
78 exact_norms.append(torch.linalg.vector_norm(u, dim=1).detach())
79 held_norms.append(torch.linalg.vector_norm(applied, dim=1).detach())
80 errors.append(torch.linalg.vector_norm(u-applied, dim=1).detach())
81 events.append(ev.detach())
82 if collect and events:
83 self.last_signature = {
84 'mean_V': float(torch.cat(vs).mean()),
85 'mean_exact_u_norm': float(torch.cat(exact_norms).mean()),
86 'mean_applied_u_norm': float(torch.cat(held_norms).mean()),
87 'mean_hold_error': float(torch.cat(errors).mean()),
88 'event_rate_per_sequence_step': float(torch.stack(events).float().mean()),
89 'max_hold_error': float(torch.cat(errors).max()),
90 }
91 return self.head(h)
92
93
94def train_one(seed, cfg, mode, collect=False, epochs=18):
95 torch.manual_seed(seed); np.random.seed(seed)
96 ds = bench.get_dataset('dynamics', seed, n_train=400, n_test=400)
97 dev = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
98 try:
99 net = PhaseGRU(mode=mode, k=cfg.get('k', 1.0), delta=cfg.get('delta', .05)).to(dev)
100 opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
101 xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev).view(-1,1)
102 for _ in range(epochs):
103 net.train()
104 perm = torch.randperm(len(xtr), device=dev)
105 for i in range(0, len(xtr), 128):
106 q = perm[i:i+128]
107 loss = ((net(xtr[q]) - ytr[q]) ** 2).mean()
108 opt.zero_grad(); loss.backward(); opt.step()
109 net.eval()
110 with torch.no_grad():
111 pred = net(ds['xte'].to(dev), collect=collect)
112 metric = float(((pred - ds['yte'].to(dev).view(-1,1))**2).mean())
113 return metric, dict(net.last_signature)
114 except Exception:
115 if dev.type == 'cuda':
116 torch.cuda.empty_cache()
117 os.environ['CUDA_VISIBLE_DEVICES'] = ''
118 return train_one(seed, cfg, mode, collect, epochs)
119 raise
120
121
122def evaluate_cfg(cfg, mode, collect=False):
123 vals, sigs = [], []
124 for s in SEEDS:
125 v, sig = train_one(s, cfg, mode, collect=collect)
126 vals.append(v); sigs.append(sig)
127 return {'mean': float(np.mean(vals)), 'std': float(np.std(vals)),
128 'per_seed': vals, 'n': len(vals), 'signatures': sigs}
129
130
131def main():
132 check = math_check()
133 # Baseline sweep covers every lr used by the idea arm.
134 base_sweep = []
135 for cfg in GRID:
136 r = evaluate_cfg(cfg, 'baseline')
137 base_sweep.append({'cfg': cfg, 'mean': r['mean'], 'std': r['std'], 'per_seed': r['per_seed']})
138 best = min(base_sweep, key=lambda q: q['mean'])
139 base_full = next(evaluate_cfg(best['cfg'], 'baseline') for _ in [0])
140 idea_grid = [dict(best['cfg'], k=1.0, delta=d) for d in (.02, .05, .12)]
141 idea_runs = []
142 for cfg in idea_grid:
143 r = evaluate_cfg(cfg, 'event', collect=True)
144 idea_runs.append({'cfg': cfg, **r})
145 ibest = min(idea_runs, key=lambda q: q['mean'])
146 diffs = [a-b for a,b in zip(ibest['per_seed'], base_full['per_seed'])]
147 p = bench.permutation_pvalue(diffs)
148 sig = {'prediction': 'event rotation should preserve pair norms and bound hold error by delta',
149 'observed': ibest['signatures'],
150 'mean_hold_error': float(np.mean([s['mean_hold_error'] for s in ibest['signatures']])),
151 'delta': ibest['cfg']['delta'],
152 'confirmed': all(s['max_hold_error'] <= ibest['cfg']['delta'] + 1e-5 for s in ibest['signatures'])}
153 base_block = {'best_cfg': best['cfg'], 'sweep': base_sweep, 'full': base_full}
154 idea_block = {'best_cfg': ibest['cfg'], 'sweep': [{'cfg': x['cfg'], 'mean': x['mean'], 'std': x['std']} for x in idea_runs], 'full': ibest}
155 report = bench.make_report('dynamics', 'rnn_small_phase_gru', base_block, ibest, extra=sig)
156 report['math_check'] = check
157 report['custom_track'] = None
158 # Preserve the complete idea sweep and explicit paired permutation result.
159 report['idea'] = idea_block
160 report['comparison']['delta_mean'] = float(np.mean(diffs))
161 report['comparison']['paired_diffs'] = diffs
162 report['comparison']['permutation_pvalue'] = p
163 Path('bench_report.json').write_text(json.dumps(report, indent=2))
164 print(json.dumps(report, indent=2))
165
166if __name__ == '__main__':
167 main()