import json, sys from pathlib import Path import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') import bench SEEDS = tuple(range(8)) GRID = [ {'lr': 0.001, 'weight_decay': 0.0}, {'lr': 0.003, 'weight_decay': 0.0}, {'lr': 0.006, 'weight_decay': 0.0}, ] def control(z, k): r = z.mean(1, keepdim=True) return (2.0 * k / z.shape[1]) * torch.imag(z * torch.conj(r)) def math_check(seed=123): rng = np.random.default_rng(seed) th = torch.tensor(rng.normal(0, .22, 24), dtype=torch.float64, requires_grad=True) z = torch.exp(1j * th) V = torch.abs(z.mean()) ** 2 grad = torch.autograd.grad(V, th)[0] dt = 1e-7 th2 = th.detach() + dt * (-grad.detach()) V2 = torch.abs(torch.exp(1j * th2).mean()) ** 2 observed = float((V2 - V.detach()) / dt) predicted = float(-(grad * grad).sum()) ratio = observed / predicted return {'Vdot_observed': observed, 'Vdot_predicted': predicted, 'ratio': ratio, 'passed': abs(ratio - 1) < 1e-4} class PhaseGRU(nn.Module): """The same GRUCell architecture in both arms; event rotation is the only change.""" def __init__(self, hidden=64, mode='baseline', k=1.0, delta=0.05): super().__init__() if hidden % 2: raise ValueError('hidden must be even') self.cell = nn.GRUCell(3, hidden) self.head = nn.Linear(hidden, 1) self.mode, self.k, self.delta = mode, k, delta self.last_signature = {} def forward(self, x, collect=False): b = x.shape[0] h = x.new_zeros(b, self.cell.hidden_size) held = x.new_zeros(b, self.cell.hidden_size // 2) vs, exact_norms, held_norms, events, errors = [], [], [], [], [] for t in range(x.shape[1] // 3): h = self.cell(x[:, 3*t:3*t+3], h) p = h.view(b, -1, 2) norm = torch.sqrt((p*p).sum(-1) + 1e-8) z = torch.complex(p[..., 0], p[..., 1]) / norm u = control(z, self.k) if self.mode == 'continuous': applied = u ev = torch.ones(b, device=x.device, dtype=torch.bool) elif self.mode == 'event': ev = torch.linalg.vector_norm(u - held, dim=1) >= self.delta applied = torch.where(ev[:, None], u, held) held = applied else: applied = torch.zeros_like(u) ev = torch.zeros(b, device=x.device, dtype=torch.bool) angle = applied.unsqueeze(-1) rot = torch.stack((p[..., 0]*torch.cos(angle[..., 0]) - p[..., 1]*torch.sin(angle[..., 0]), p[..., 0]*torch.sin(angle[..., 0]) + p[..., 1]*torch.cos(angle[..., 0])), -1) h = rot.reshape(b, -1) if collect: r = z.mean(1) vs.append((r.abs()**2).detach()) exact_norms.append(torch.linalg.vector_norm(u, dim=1).detach()) held_norms.append(torch.linalg.vector_norm(applied, dim=1).detach()) errors.append(torch.linalg.vector_norm(u-applied, dim=1).detach()) events.append(ev.detach()) if collect and events: self.last_signature = { 'mean_V': float(torch.cat(vs).mean()), 'mean_exact_u_norm': float(torch.cat(exact_norms).mean()), 'mean_applied_u_norm': float(torch.cat(held_norms).mean()), 'mean_hold_error': float(torch.cat(errors).mean()), 'event_rate_per_sequence_step': float(torch.stack(events).float().mean()), 'max_hold_error': float(torch.cat(errors).max()), } return self.head(h) def train_one(seed, cfg, mode, collect=False, epochs=18): torch.manual_seed(seed); np.random.seed(seed) ds = bench.get_dataset('dynamics', seed, n_train=400, n_test=400) dev = torch.device('cuda' if torch.cuda.is_available() else 'cpu') try: net = PhaseGRU(mode=mode, k=cfg.get('k', 1.0), delta=cfg.get('delta', .05)).to(dev) opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay']) xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev).view(-1,1) for _ in range(epochs): net.train() perm = torch.randperm(len(xtr), device=dev) for i in range(0, len(xtr), 128): q = perm[i:i+128] loss = ((net(xtr[q]) - ytr[q]) ** 2).mean() opt.zero_grad(); loss.backward(); opt.step() net.eval() with torch.no_grad(): pred = net(ds['xte'].to(dev), collect=collect) metric = float(((pred - ds['yte'].to(dev).view(-1,1))**2).mean()) return metric, dict(net.last_signature) except Exception: if dev.type == 'cuda': torch.cuda.empty_cache() os.environ['CUDA_VISIBLE_DEVICES'] = '' return train_one(seed, cfg, mode, collect, epochs) raise def evaluate_cfg(cfg, mode, collect=False): vals, sigs = [], [] for s in SEEDS: v, sig = train_one(s, cfg, mode, collect=collect) vals.append(v); sigs.append(sig) return {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals), 'signatures': sigs} def main(): check = math_check() # Baseline sweep covers every lr used by the idea arm. base_sweep = [] for cfg in GRID: r = evaluate_cfg(cfg, 'baseline') base_sweep.append({'cfg': cfg, 'mean': r['mean'], 'std': r['std'], 'per_seed': r['per_seed']}) best = min(base_sweep, key=lambda q: q['mean']) base_full = next(evaluate_cfg(best['cfg'], 'baseline') for _ in [0]) idea_grid = [dict(best['cfg'], k=1.0, delta=d) for d in (.02, .05, .12)] idea_runs = [] for cfg in idea_grid: r = evaluate_cfg(cfg, 'event', collect=True) idea_runs.append({'cfg': cfg, **r}) ibest = min(idea_runs, key=lambda q: q['mean']) diffs = [a-b for a,b in zip(ibest['per_seed'], base_full['per_seed'])] p = bench.permutation_pvalue(diffs) sig = {'prediction': 'event rotation should preserve pair norms and bound hold error by delta', 'observed': ibest['signatures'], 'mean_hold_error': float(np.mean([s['mean_hold_error'] for s in ibest['signatures']])), 'delta': ibest['cfg']['delta'], 'confirmed': all(s['max_hold_error'] <= ibest['cfg']['delta'] + 1e-5 for s in ibest['signatures'])} base_block = {'best_cfg': best['cfg'], 'sweep': base_sweep, 'full': base_full} idea_block = {'best_cfg': ibest['cfg'], 'sweep': [{'cfg': x['cfg'], 'mean': x['mean'], 'std': x['std']} for x in idea_runs], 'full': ibest} report = bench.make_report('dynamics', 'rnn_small_phase_gru', base_block, ibest, extra=sig) report['math_check'] = check report['custom_track'] = None # Preserve the complete idea sweep and explicit paired permutation result. report['idea'] = idea_block report['comparison']['delta_mean'] = float(np.mean(diffs)) report['comparison']['paired_diffs'] = diffs report['comparison']['permutation_pvalue'] = p Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()