Geometrically Attracting Random Recurrent Layer / bench_run.py

Failed on benchmark

Raw ⬇ ZIP
  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()