Schäffer-Covariant Isometric Recurrent Layer / stage2_bench.py
Mechanism confirmed, baseline not beaten
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()