Schäffer-Covariant Isometric Recurrent Layer / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5import sys
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10GRID = [{'lr': 1e-3, 'epochs': 8}, {'lr': 3e-3, 'epochs': 8}, {'lr': 1e-2, 'epochs': 8}]
 11
 12class SpectralLinearRNN(nn.Module):
 13    def __init__(self, input_dim=3, out_dim=1, hidden=24):
 14        super().__init__()
 15        self.d = hidden
 16        self.raw = nn.Parameter(torch.randn(hidden, hidden) * .08)
 17        self.inp = nn.Linear(input_dim, hidden)
 18        self.head = nn.Linear(hidden, out_dim)
 19    def transition(self):
 20        # differentiable contraction normalization, standard baseline mechanism
 21        return self.raw / torch.clamp(torch.linalg.matrix_norm(self.raw, 2), min=1.0)
 22    def forward(self, x):
 23        x = x.reshape(x.shape[0], -1, 3)
 24        h = torch.zeros(x.shape[0], self.d, device=x.device, dtype=x.dtype)
 25        X = self.transition()
 26        for t in range(x.shape[1]):
 27            h = h @ X.T + self.inp(x[:, t])
 28        return self.head(h)
 29
 30class SchafferLiftRNN(nn.Module):
 31    def __init__(self, input_dim=3, out_dim=1, hidden=24, K=8):
 32        super().__init__()
 33        self.d, self.K = hidden, K
 34        self.raw = nn.Parameter(torch.randn(hidden, hidden) * .08)
 35        self.inp = nn.Linear(input_dim, hidden)
 36        self.head = nn.Linear(hidden, out_dim)
 37    def transition(self):
 38        return self.raw / torch.clamp(torch.linalg.matrix_norm(self.raw, 2), min=1.0)
 39    def defect(self, X):
 40        I = torch.eye(self.d, device=X.device, dtype=X.dtype)
 41        A = (I - X.T @ X + (I - X.T @ X).T) / 2
 42        w, U = torch.linalg.eigh(A)
 43        return (U * torch.sqrt(torch.clamp(w, min=1e-7))) @ U.T
 44    def forward(self, x):
 45        x = x.reshape(x.shape[0], -1, 3)
 46        B = x.shape[0]; dev=x.device; dtype=x.dtype
 47        h = torch.zeros(B, self.d, device=dev, dtype=dtype)
 48        q = torch.zeros(B, self.K, self.d, device=dev, dtype=dtype)
 49        X = self.transition(); D = self.defect(X)
 50        for t in range(x.shape[1]):
 51            # exact finite Schaeffer update, with external input entering active state
 52            h_new = h @ X.T + self.inp(x[:, t])
 53            q_new = torch.cat([(h @ D.T).unsqueeze(1), q[:, :-1]], dim=1)
 54            h, q = h_new, q_new
 55        return self.head(h)
 56
 57def seed_all(seed):
 58    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 59
 60def run_one(kind, cfg, seed):
 61    seed_all(seed)
 62    ds = get_dataset('dynamics', seed, n_train=400, n_test=200)
 63    model = SpectralLinearRNN() if kind == 'baseline' else SchafferLiftRNN()
 64    net, metric, hist = train_model(model, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128, log=lambda *_: None)
 65    if net is None: raise RuntimeError('training failed')
 66    return float(metric), net
 67
 68def train_value(kind, cfg, seed):
 69    return run_one(kind, cfg, seed)[0]
 70
 71def signature():
 72    metric, net = run_one('idea', GRID[1], 0)
 73    net.eval(); X = net.transition().detach(); D = net.defect(X).detach()
 74    d=net.d; K=net.K; n=(K+1)*d
 75    Y=torch.zeros(n,n,device=X.device); Y[:d,:d]=X; Y[d:2*d,:d]=D
 76    for j in range(1,K): Y[(j+1)*d:(j+2)*d,j*d:(j+1)*d]=torch.eye(d,device=X.device)
 77    z=torch.randn(n,device=X.device); z[d:]=0; z0=z.clone(); z=Y@z
 78    one_err=float(abs(z.norm().item()**2-z0.norm().item()**2)/(z0.norm().item()**2+1e-12))
 79    # Compare the trained transition with an uncorrected Euler-like repeated update.
 80    e=torch.randn(d,device=X.device); e0=e.clone(); z=torch.randn(n,device=X.device); z[d:]=0; z0=z.clone()
 81    for _ in range(min(4, K-1)): e=X@e; z=Y@z
 82    observed_lift=float(z.norm()/z0.norm()); observed_base=float(e.norm()/e0.norm())
 83    return {'prediction':'Schaeffer active-subspace update preserves total energy before memory truncation; baseline contraction decays',
 84            'observed_one_step_relative_energy_error':one_err,
 85            'observed_lift_preboundary_norm_ratio':observed_lift,
 86            'observed_baseline_same_horizon_norm_ratio':observed_base,
 87            'trained_test_mse':metric,
 88            'confirmed': bool(one_err < 1e-4 and abs(observed_lift-1) < 1e-3 and observed_base < .999)}
 89
 90def main():
 91    # Canonical harness sweep on four seeds, with the same union of settings on both sides.
 92    def mk(cfg): return lambda seed: train_value('baseline', cfg, int(seed))
 93    tuned = sweep_baseline(mk, GRID, seeds=(0,1,2,3))
 94    base_sweep=[]
 95    for cfg in GRID:
 96        r=evaluate(mk(cfg), seeds=SEEDS); base_sweep.append({'cfg':cfg, **r})
 97    best=min(base_sweep, key=lambda r:r['mean'])
 98    base={'best_cfg':best['cfg'], 'sweep':base_sweep, 'harness_tuning':tuned,
 99          'full':evaluate(mk(best['cfg']), seeds=SEEDS)}
100    idea_sweep=[]
101    for cfg in GRID:
102        r=evaluate(lambda seed, c=cfg: train_value('idea', c, int(seed)), seeds=SEEDS)
103        idea_sweep.append({'cfg':cfg, **r})
104    ibest=min(idea_sweep, key=lambda r:r['mean'])
105    idea={k:ibest[k] for k in ('cfg','per_seed','mean','std','n') if k in ibest}
106    rep=make_report('dynamics','schaffer_lift_vs_spectral_linear_rnn',base,idea,
107                    {'idea_sweep':idea_sweep, 'mechanism_signature':signature()})
108    rep['custom_track']=None
109    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
110    print(json.dumps(rep,indent=2))
111
112if __name__ == '__main__': main()