Forward-Intersection Spectral Latent Dynamics / bench_runner.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, make_report, sweep_baseline, evaluate
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = (0, 1, 2, 3)
 11EPOCHS = 10
 12BATCH = 128
 13NTRAIN, NTEST = 1200, 400
 14
 15
 16def ds_for(seed):
 17    return get_dataset('dynamics', int(seed), n_train=NTRAIN, n_test=NTEST)
 18
 19
 20def base_run(cfg, seed):
 21    torch.manual_seed(seed); np.random.seed(seed)
 22    d = ds_for(seed)
 23    net = make_model('rnn_small', d['input_shape'], d['out_dim'])
 24    _, metric, _ = train_model(net, d, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH,
 25                               weight_decay=cfg.get('weight_decay', 0.0), log=lambda *_: None)
 26    return float(metric)
 27
 28
 29class ForwardIntersectionRNN(nn.Module):
 30    """rnn_small with a detached forward-compatible hidden-state projection."""
 31    def __init__(self, out_dim, hidden=64, tau=0.15, levels=1):
 32        super().__init__()
 33        self.rnn = nn.GRU(3, hidden, batch_first=True)
 34        self.head = nn.Linear(hidden, out_dim)
 35        self.tau, self.levels = float(tau), int(levels)
 36        self.register_buffer('Q', torch.eye(hidden))
 37        self.register_buffer('A_est', torch.eye(hidden))
 38        self.register_buffer('active', torch.tensor(0, dtype=torch.int64))
 39
 40    def forward(self, x):
 41        seq = x.view(x.shape[0], -1, 3)
 42        out, h = self.rnn(seq)
 43        q = self.Q
 44        if int(self.active.item()):
 45            out = out @ q @ q.T
 46            h = h @ q @ q.T
 47        return self.head(h[-1])
 48
 49    @torch.no_grad()
 50    def refresh(self, x):
 51        """Estimate A on consecutive hidden states and retain compatible directions."""
 52        was_training = self.training
 53        self.eval()
 54        seq = x.view(x.shape[0], -1, 3)
 55        out, _ = self.rnn(seq)
 56        # Consecutive hidden states in a batch provide snapshot pairs.
 57        X, Y = out[:, :-1, :].reshape(-1, out.shape[-1]), out[:, 1:, :].reshape(-1, out.shape[-1])
 58        if X.shape[0] < 4:
 59            return
 60        A = (torch.linalg.lstsq(X, Y).solution).T
 61        # Principal-angle compatibility of range(I) and range(A) reduces to
 62        # singular directions of A; retain directions with singular values near 1.
 63        U, s, _ = torch.linalg.svd(A)
 64        scale = torch.clamp(s.max(), min=1e-6)
 65        # A direction is forward-supported when its normalized image is not
 66        # strongly collapsed; this is a noise-robust finite-dimensional proxy.
 67        keep = s / scale >= (1.0 - self.tau)
 68        if int(keep.sum()) < 1:
 69            keep[torch.argmax(s)] = True
 70        q = U[:, keep]
 71        # Additional levels repeatedly apply the same compatibility test.
 72        for _ in range(max(0, self.levels - 1)):
 73            B = q.T @ A @ q
 74            u2, s2, _ = torch.linalg.svd(B)
 75            k2 = s2 / torch.clamp(s2.max(), min=1e-6) >= (1.0 - self.tau)
 76            if int(k2.sum()) == 0: break
 77            q = q @ u2[:, k2]
 78        self.Q.zero_()
 79        self.Q[:q.shape[0], :q.shape[1]] = q
 80        # Q is stored padded; active dimension records selected rank.
 81        self.A_est.zero_(); self.A_est[:A.shape[0], :A.shape[1]] = A
 82        self.active.fill_(1)
 83        if was_training: self.train()
 84
 85
 86def idea_run(cfg, seed, return_model=False):
 87    torch.manual_seed(seed); np.random.seed(seed)
 88    d = ds_for(seed)
 89    net = ForwardIntersectionRNN(d['out_dim'], hidden=64, tau=cfg['tau'], levels=cfg['levels'])
 90    # Identical Adam/MSE training budget; refresh once from training snapshots
 91    # after optimization, so the intervention is used in the evaluated system.
 92    opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg.get('weight_decay', 0.0))
 93    lossf = nn.MSELoss()
 94    xtr, ytr = d['xtr'], d['ytr']
 95    for _ in range(EPOCHS):
 96        net.train(); perm = torch.randperm(len(xtr))
 97        for i in range(0, len(xtr), BATCH):
 98            ix = perm[i:i+BATCH]
 99            loss = lossf(net(xtr[ix]), ytr[ix])
100            opt.zero_grad(); loss.backward(); opt.step()
101    net.refresh(xtr)
102    net.eval()
103    with torch.no_grad(): metric = float(lossf(net(d['xte']), d['yte']))
104    if return_model: return metric, net, d
105    return metric
106
107
108def signature(cfg, seeds=(0, 1, 2, 3)):
109    rows=[]
110    for s in seeds:
111        metric, net, d = idea_run(cfg, s, True)
112        with torch.no_grad():
113            seq=d['xte'][:128].view(-1,8,3); h,_=net.rnn(seq)
114            before=h.reshape(-1,64)
115            after=(before @ net.Q @ net.Q.T)
116            raw_next=before[:,1:] if False else before
117            residual=float(torch.mean((after-before)**2))
118            rank=int(net.active.item() and torch.linalg.matrix_rank(net.Q).item() or 64)
119        eig=np.linalg.eigvals(net.A_est.cpu().numpy())
120        rows.append({'seed':s,'metric':metric,'projection_mse':residual,'rank':rank,'raw_spectral_radius':float(np.max(np.abs(eig)))})
121    return rows
122
123
124def main():
125    # lr union is shared by both sides; baseline central knob includes weight decay.
126    grid=[{'lr':1e-3,'weight_decay':0.0},{'lr':3e-3,'weight_decay':0.0},
127          {'lr':1e-2,'weight_decay':0.0},{'lr':3e-3,'weight_decay':1e-4}]
128    base=sweep_baseline(lambda cfg: lambda seed: base_run(cfg, seed), grid, seeds=SWEEP_SEEDS)
129    idea_grid=[{'lr':base['best_cfg']['lr'],'weight_decay':base['best_cfg'].get('weight_decay',0.0),'tau':t,'levels':1} for t in (0.10,0.15,0.25)]
130    # Baseline was evaluated at every lr/weight-decay in the union above.
131    best_idea_cfg=min(idea_grid, key=lambda c: np.mean([idea_run(c,s) for s in SWEEP_SEEDS]))
132    idea=evaluate(lambda s: idea_run(best_idea_cfg,s), seeds=SEEDS)
133    sigrows=signature(best_idea_cfg)
134    sig={'prediction':'compatible projection reduces unsupported hidden transition energy/rank without changing task architecture',
135         'observed':sigrows,
136         'mean_projection_mse':float(np.mean([r['projection_mse'] for r in sigrows])),
137         'mean_rank':float(np.mean([r['rank'] for r in sigrows])),
138         'confirmed':bool(np.mean([r['projection_mse'] for r in sigrows])>1e-8 and np.mean([r['rank'] for r in sigrows])<64)}
139    report=make_report('dynamics','rnn_small',base,idea,{'mechanism_signature':sig,'custom_track':None,'idea_cfg':best_idea_cfg,'idea_grid':idea_grid})
140    print(json.dumps(report, indent=2))
141
142if __name__=='__main__': main()