import json, sys from pathlib import Path import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report from isometric_mixer import IsometricButterflyMixer SEEDS = tuple(range(8)) SWEEP_SEEDS = tuple(range(4)) GRID = [{'lr': lr, 'epochs': 6, 'weight_decay': 0.0} for lr in (0.001, 0.003, 0.006)] class SequenceMixerEncoder(nn.Module): """Matched sequence model: only token-mixing operator differs.""" def __init__(self, n_tokens, out_dim, width=64, depth=2, isometric=False): super().__init__() self.n_tokens, self.isometric = n_tokens, isometric self.inp = nn.Linear(1, width) self.pos = nn.Parameter(torch.zeros(1, n_tokens, width)) nn.init.normal_(self.pos, std=0.02) self.mixers = nn.ModuleList() self.norms = nn.ModuleList() self.ffns = nn.ModuleList() for _ in range(depth): self.mixers.append(IsometricButterflyMixer(n_tokens) if isometric else nn.Linear(n_tokens, n_tokens)) self.norms.append(nn.LayerNorm(width)) self.ffns.append(nn.Sequential(nn.Linear(width, 128), nn.GELU(), nn.Linear(128, width))) self.head = nn.Linear(n_tokens * width, out_dim) def forward(self, x): h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]] for mixer, norm, ffn in zip(self.mixers, self.norms, self.ffns): if self.isometric: mixed = mixer(h) else: mixed = mixer(h.transpose(-1, -2)).transpose(-1, -2) h = h + mixed h = h + ffn(norm(h)) return self.head(h.reshape(h.shape[0], -1)) def train_one(seed, cfg, idea=False, capture=False): torch.manual_seed(seed); np.random.seed(seed) ds = get_dataset('sequence', seed, n_train=400, n_test=400) model = SequenceMixerEncoder(ds['input_shape'][0], ds['out_dim'], isometric=idea) trained, metric, _history = train_model(model, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128, weight_decay=cfg.get('weight_decay', 0.0), log=lambda *_: None) if trained is None or metric is None: return float('inf') if not capture else {'metric': float('inf'), 'norm_ratio': float('nan'), 'inverse_error': float('nan')} model = trained if not capture: return float(metric) device = next(model.parameters()).device with torch.no_grad(): x = ds['xte'][:32].to(device) base = model.inp(x.unsqueeze(-1)) if idea: z = model.mixers[0](base) back = model.mixers[0].adjoint(z) norm_ratio = float(z.norm() / base.norm()) inverse_error = float((back - base).abs().max()) else: z = model.mixers[0](base.transpose(-1, -2)).transpose(-1, -2) norm_ratio = float(z.norm() / base.norm()) inverse_error = None return {'metric': float(metric), 'norm_ratio': norm_ratio, 'inverse_error': inverse_error} def math_check(): vals=[] for n in (2,4,8,16,32): torch.manual_seed(100+n); m=IsometricButterflyMixer(n) x=torch.randn(7,n,3); y=m(x); rec=m.adjoint(y) vals.append({'n': n, 'norm_ratio': float(y.norm()/x.norm()), 'inverse_max_error': float((rec-x).abs().max())}) return vals def main(): def base_fn(cfg): return lambda seed: train_one(seed, cfg, False) baseline = sweep_baseline(base_fn, GRID, seeds=SWEEP_SEEDS) idea_runs=[] for cfg in GRID: r=evaluate(lambda seed, c=cfg: train_one(seed, c, True), seeds=SEEDS) idea_runs.append({'cfg':cfg, 'result':r}) best=min(idea_runs, key=lambda q:q['result']['mean']) idea=best['result'] sig=train_one(0, best['cfg'], True, capture=True) base_sig=train_one(0, best['cfg'], False, capture=True) report=make_report('sequence','transformer_tiny',baseline,idea,extra={ 'idea_sweep': idea_runs, 'observed_best_cfg': best['cfg'], 'math_check': math_check(), 'mechanism_signature': { 'prediction': 'trained isometric token mixer preserves feature-token Euclidean norm and has exact adjoint inverse', 'predicted_norm_ratio': 1.0, 'observed_norm_ratio_seed0': sig['norm_ratio'], 'observed_inverse_max_error_seed0': sig['inverse_error'], 'baseline_observed_norm_ratio_seed0': base_sig['norm_ratio'], 'confirmed': abs(sig['norm_ratio']-1.0)<1e-5 and sig['inverse_error']<2e-5, 'trained_model_metric_seed0': sig['metric'] } }) Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()