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()