Isometric tensor-network token mixer / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import json, sys
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, train_model, evaluate, sweep_baseline, make_report
9from isometric_mixer import IsometricButterflyMixer
10
11SEEDS = tuple(range(8))
12SWEEP_SEEDS = tuple(range(4))
13GRID = [{'lr': lr, 'epochs': 6, 'weight_decay': 0.0} for lr in (0.001, 0.003, 0.006)]
14
15
16class SequenceMixerEncoder(nn.Module):
17 """Matched sequence model: only token-mixing operator differs."""
18 def __init__(self, n_tokens, out_dim, width=64, depth=2, isometric=False):
19 super().__init__()
20 self.n_tokens, self.isometric = n_tokens, isometric
21 self.inp = nn.Linear(1, width)
22 self.pos = nn.Parameter(torch.zeros(1, n_tokens, width))
23 nn.init.normal_(self.pos, std=0.02)
24 self.mixers = nn.ModuleList()
25 self.norms = nn.ModuleList()
26 self.ffns = nn.ModuleList()
27 for _ in range(depth):
28 self.mixers.append(IsometricButterflyMixer(n_tokens) if isometric else nn.Linear(n_tokens, n_tokens))
29 self.norms.append(nn.LayerNorm(width))
30 self.ffns.append(nn.Sequential(nn.Linear(width, 128), nn.GELU(), nn.Linear(128, width)))
31 self.head = nn.Linear(n_tokens * width, out_dim)
32
33 def forward(self, x):
34 h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
35 for mixer, norm, ffn in zip(self.mixers, self.norms, self.ffns):
36 if self.isometric:
37 mixed = mixer(h)
38 else:
39 mixed = mixer(h.transpose(-1, -2)).transpose(-1, -2)
40 h = h + mixed
41 h = h + ffn(norm(h))
42 return self.head(h.reshape(h.shape[0], -1))
43
44
45def train_one(seed, cfg, idea=False, capture=False):
46 torch.manual_seed(seed); np.random.seed(seed)
47 ds = get_dataset('sequence', seed, n_train=400, n_test=400)
48 model = SequenceMixerEncoder(ds['input_shape'][0], ds['out_dim'], isometric=idea)
49 trained, metric, _history = train_model(model, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128,
50 weight_decay=cfg.get('weight_decay', 0.0), log=lambda *_: None)
51 if trained is None or metric is None:
52 return float('inf') if not capture else {'metric': float('inf'), 'norm_ratio': float('nan'), 'inverse_error': float('nan')}
53 model = trained
54 if not capture:
55 return float(metric)
56 device = next(model.parameters()).device
57 with torch.no_grad():
58 x = ds['xte'][:32].to(device)
59 base = model.inp(x.unsqueeze(-1))
60 if idea:
61 z = model.mixers[0](base)
62 back = model.mixers[0].adjoint(z)
63 norm_ratio = float(z.norm() / base.norm())
64 inverse_error = float((back - base).abs().max())
65 else:
66 z = model.mixers[0](base.transpose(-1, -2)).transpose(-1, -2)
67 norm_ratio = float(z.norm() / base.norm())
68 inverse_error = None
69 return {'metric': float(metric), 'norm_ratio': norm_ratio, 'inverse_error': inverse_error}
70
71
72def math_check():
73 vals=[]
74 for n in (2,4,8,16,32):
75 torch.manual_seed(100+n); m=IsometricButterflyMixer(n)
76 x=torch.randn(7,n,3); y=m(x); rec=m.adjoint(y)
77 vals.append({'n': n, 'norm_ratio': float(y.norm()/x.norm()),
78 'inverse_max_error': float((rec-x).abs().max())})
79 return vals
80
81
82def main():
83 def base_fn(cfg): return lambda seed: train_one(seed, cfg, False)
84 baseline = sweep_baseline(base_fn, GRID, seeds=SWEEP_SEEDS)
85 idea_runs=[]
86 for cfg in GRID:
87 r=evaluate(lambda seed, c=cfg: train_one(seed, c, True), seeds=SEEDS)
88 idea_runs.append({'cfg':cfg, 'result':r})
89 best=min(idea_runs, key=lambda q:q['result']['mean'])
90 idea=best['result']
91 sig=train_one(0, best['cfg'], True, capture=True)
92 base_sig=train_one(0, best['cfg'], False, capture=True)
93 report=make_report('sequence','transformer_tiny',baseline,idea,extra={
94 'idea_sweep': idea_runs,
95 'observed_best_cfg': best['cfg'],
96 'math_check': math_check(),
97 'mechanism_signature': {
98 'prediction': 'trained isometric token mixer preserves feature-token Euclidean norm and has exact adjoint inverse',
99 'predicted_norm_ratio': 1.0,
100 'observed_norm_ratio_seed0': sig['norm_ratio'],
101 'observed_inverse_max_error_seed0': sig['inverse_error'],
102 'baseline_observed_norm_ratio_seed0': base_sig['norm_ratio'],
103 'confirmed': abs(sig['norm_ratio']-1.0)<1e-5 and sig['inverse_error']<2e-5,
104 'trained_model_metric_seed0': sig['metric']
105 }
106 })
107 Path('bench_report.json').write_text(json.dumps(report, indent=2))
108 print(json.dumps(report, indent=2))
109
110if __name__ == '__main__': main()