Isometric tensor-network token mixer / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()