Parabolic Riesz Feature Preconditioner / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
9
10NTRAIN, NTEST = 400, 200
11EPOCHS = 10
12BATCH = 128
13GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 6e-3}]
14
15
16def riesz_matrix(n, lam=1.0, eps=1e-5):
17 D = np.zeros((n, n), dtype=np.float32)
18 for i in range(n):
19 D[i, i] = -1.0
20 D[i, (i + 1) % n] = 1.0
21 L = D.T @ D
22 H = L + eps * np.eye(n, dtype=np.float32)
23 R = lam * D @ np.linalg.inv(np.eye(n, dtype=np.float32) + lam * lam * H)
24 return torch.tensor(R), torch.tensor(D)
25
26
27class RieszTransformer(nn.Module):
28 def __init__(self, win, out_dim, lam=1.0):
29 super().__init__()
30 d = 64
31 self.inp = nn.Linear(1, d)
32 self.pos = nn.Parameter(torch.zeros(1, win, d))
33 nn.init.normal_(self.pos, std=.02)
34 R, D = riesz_matrix(win, lam)
35 self.register_buffer('R', R)
36 self.register_buffer('D', D)
37 self.branch = nn.Linear(d, d, bias=False)
38 self.g = nn.Parameter(torch.tensor(0.1))
39 layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=128,
40 batch_first=True, dropout=0.0)
41 self.enc = nn.TransformerEncoder(layer, 2)
42 self.head = nn.Linear(win * d, out_dim)
43
44 def forward(self, x):
45 h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
46 z = torch.einsum('ij,bjd->bid', self.R, h)
47 z = z / (z.square().mean(dim=(1, 2), keepdim=True).sqrt() + 1e-5)
48 h = h + self.g * self.branch(z)
49 return self.head(self.enc(h).reshape(x.shape[0], -1))
50
51
52def seed_all(seed):
53 random.seed(seed)
54 np.random.seed(seed)
55 torch.manual_seed(seed)
56 if torch.cuda.is_available():
57 torch.cuda.manual_seed_all(seed)
58
59
60def train_one(kind, cfg, seed, return_model=False):
61 seed_all(seed)
62 ds = get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST)
63 if kind == 'baseline':
64 model = make_model('transformer_tiny', ds['input_shape'], ds['out_dim'])
65 else:
66 model = RieszTransformer(ds['input_shape'][0], ds['out_dim'])
67 net, metric, hist = train_model(model, ds, epochs=EPOCHS,
68 lr=cfg['lr'], batch=BATCH, log=lambda _: None)
69 if net is None:
70 raise RuntimeError('training failed')
71 if return_model:
72 return float(metric), net, ds
73 return float(metric)
74
75
76def signature():
77 vals = []
78 for seed in range(8):
79 b, bm, ds = train_one('baseline', {'lr': 3e-3}, seed, True)
80 r, rm, _ = train_one('idea', {'lr': 3e-3}, seed, True)
81 devb = next(bm.parameters()).device
82 devi = next(rm.parameters()).device
83 xb_in = ds['xte'][:64].to(devb)
84 ri_in = ds['xte'][:64].to(devi)
85 noise_b = torch.randn_like(xb_in) * 0.10
86 noise_i = noise_b.to(devi)
87 with torch.no_grad():
88 xb = bm(xb_in); xbn = bm(xb_in + noise_b)
89 ri = rm(ri_in); rin = rm(ri_in + noise_i)
90 vals.append({'baseline_output_noise_ratio': float((xbn-xb).norm()/(xb.norm()+1e-8)),
91 'idea_output_noise_ratio': float((rin-ri).norm()/(ri.norm()+1e-8)),
92 'baseline_metric': b, 'idea_metric': r})
93 br = np.mean([v['baseline_output_noise_ratio'] for v in vals])
94 ir = np.mean([v['idea_output_noise_ratio'] for v in vals])
95 return {'mean_baseline_output_noise_ratio': float(br),
96 'mean_idea_output_noise_ratio': float(ir),
97 'predicted': 'idea should attenuate feature perturbation',
98 'confirmed': bool(ir < br), 'per_seed': vals}
99
100
101def main():
102 seed_all(425)
103 def baseline_fn(cfg):
104 return lambda seed: train_one('baseline', cfg, seed)
105 base = sweep_baseline(baseline_fn, GRID, seeds=(0, 1, 2, 3))
106 # Full eight-seed idea run at best baseline lr and two nearby union-parity settings.
107 idea_runs = []
108 for cfg in GRID:
109 idea_runs.append((cfg, evaluate(lambda seed, c=cfg: train_one('idea', c, seed))))
110 best_cfg, idea_res = min(idea_runs, key=lambda z: z[1]['mean'])
111 sig = signature()
112 report = make_report('sequence', 'transformer_tiny', base, idea_res,
113 {'mechanism_signature': sig,
114 'track_match': 'sequence-level correlated multi-token forecast',
115 'idea_best_cfg': best_cfg,
116 'idea_grid': [{'cfg': c, 'full': r} for c, r in idea_runs],
117 'epochs': EPOCHS, 'n_train': NTRAIN, 'n_test': NTEST})
118 Path('bench_report.json').write_text(json.dumps(report, indent=2))
119 print(json.dumps(report, indent=2))
120
121
122if __name__ == '__main__':
123 main()