Fejer reflection accelerator for fixed-point layers / stage2_fejer_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import random
  3from pathlib import Path
  4import numpy as np
  5import torch
  6import torch.nn as nn
  7
  8import sys
  9sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 10from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
 11
 12SEEDS = tuple(range(8))
 13# Shared union: both methods are evaluated at every lr and K.
 14GRID = [
 15    {'lr': 1e-3, 'K': 3}, {'lr': 3e-3, 'K': 3}, {'lr': 5e-3, 'K': 3},
 16    {'lr': 1e-3, 'K': 7}, {'lr': 3e-3, 'K': 7}, {'lr': 5e-3, 'K': 7},
 17    {'lr': 1e-3, 'K': 15}, {'lr': 3e-3, 'K': 15}, {'lr': 5e-3, 'K': 15},
 18]
 19
 20
 21def seed_all(seed):
 22    random.seed(seed)
 23    np.random.seed(seed)
 24    torch.manual_seed(seed)
 25    if torch.cuda.is_available():
 26        torch.cuda.manual_seed_all(seed)
 27
 28
 29class FixedPointRNN(nn.Module):
 30    """Matched recurrent system with an explicit learned state map J_x.
 31
 32    baseline: s <- J_x(s), K times
 33    fejer:    z <- 2 J_x(z)-z and average s,z_1,...,z_K
 34    Both use exactly K calls to the same J_x per input token.
 35    """
 36    def __init__(self, out_dim=1, hidden=48, K=7, fejer=False):
 37        super().__init__()
 38        self.inp = nn.Linear(3, hidden)
 39        self.state = nn.Linear(hidden, hidden, bias=False)
 40        self.head = nn.Linear(hidden, out_dim)
 41        self.K = int(K)
 42        self.fejer = bool(fejer)
 43
 44    def J(self, s, drive):
 45        # A bounded contractive-style map; monotonicity is a hypothesis for NN use.
 46        return torch.tanh(drive + self.state(s))
 47
 48    def forward(self, x):
 49        seq = x.view(x.shape[0], -1, 3)
 50        s = torch.zeros(x.shape[0], self.state.in_features,
 51                        device=x.device, dtype=x.dtype)
 52        for token in seq.unbind(1):
 53            drive = self.inp(token)
 54            if self.fejer:
 55                z = s
 56                acc = s
 57                for _ in range(self.K):
 58                    z = 2.0 * self.J(z, drive) - z
 59                    acc = acc + z
 60                s = acc / (self.K + 1)
 61            else:
 62                for _ in range(self.K):
 63                    s = self.J(s, drive)
 64        return self.head(s)
 65
 66
 67def run(kind, cfg, seed, epochs=16):
 68    seed_all(seed)
 69    ds = get_dataset('dynamics', seed, n_train=400, n_test=200)
 70    net = FixedPointRNN(ds['out_dim'], K=cfg['K'], fejer=(kind == 'idea'))
 71    _, metric, _ = train_model(net, ds, epochs=epochs, lr=cfg['lr'],
 72                               batch=128, log=lambda *_: None)
 73    if metric is None:
 74        return float('nan')
 75    return float(metric)
 76
 77
 78def base_fn(cfg):
 79    return lambda seed: run('baseline', cfg, seed)
 80
 81
 82def idea_fn(cfg):
 83    return lambda seed: run('idea', cfg, seed)
 84
 85
 86def mechanism_signature():
 87    """Measure the claimed residual scaling on a trained benchmark model."""
 88    seed_all(0)
 89    ds = get_dataset('dynamics', 0, n_train=400, n_test=200)
 90    cfg = {'lr': 3e-3, 'K': 7}
 91    net = FixedPointRNN(ds['out_dim'], K=cfg['K'], fejer=True)
 92    net, _, _ = train_model(net, ds, epochs=16, lr=cfg['lr'], batch=128,
 93                            log=lambda *_: None)
 94    dev = next(net.parameters()).device
 95    x = ds['xte'][:32].to(dev)
 96    seq = x.view(x.shape[0], -1, 3)
 97    with torch.no_grad():
 98        s = torch.zeros(32, 48, device=dev)
 99        ratios, identity_errs = [], []
100        for token in seq.unbind(1):
101            drive = net.inp(token)
102            z, acc = s, s
103            for _ in range(cfg['K']):
104                z = 2 * net.J(z, drive) - z
105                acc = acc + z
106            yh = acc / (cfg['K'] + 1)
107            r = net.J(yh, drive) - yh
108            rhs = (z - s) / (2 * (cfg['K'] + 1))
109            identity_errs.append(float((r - rhs).norm(dim=1).mean().cpu()))
110            old_r = (net.J(s, drive) - s).norm(dim=1).mean()
111            new_r = r.norm(dim=1).mean()
112            ratios.append(float((new_r / old_r.clamp_min(1e-8)).cpu()))
113            s = yh
114    observed = float(np.mean(ratios))
115    predicted = 1.0 / (cfg['K'] + 1)
116    # This is a quantitative NN-scale check, not the task comparison.
117    confirmed = bool(np.isfinite(observed) and np.isfinite(predicted) and
118                     abs(observed - predicted) <= max(0.20, 2.0 * predicted))
119    return {
120        'prediction': 'Fejer residual contracts approximately as 1/(K+1) under resolvent/nonexpansive assumptions',
121        'K': cfg['K'], 'predicted_ratio': predicted,
122        'observed_trained_ratio': observed,
123        'identity_error_on_trained_model': float(max(identity_errs)),
124        'confirmed': confirmed,
125    }
126
127
128def main():
129    # Baseline sweep uses the canonical four seed tuning split; final is eight seeds.
130    baseline = sweep_baseline(base_fn, GRID)
131    idea_candidates = []
132    for cfg in GRID:
133        idea_candidates.append({'cfg': cfg, 'result': evaluate(idea_fn(cfg), seeds=SEEDS)})
134    best = min(idea_candidates, key=lambda q: q['result']['mean'])
135    report = make_report(
136        'dynamics', 'rnn_small', baseline, best['result'],
137        {'structural_match': 'controlled pendulum rollout with recurrent fixed-point state updates',
138         'shared_architecture': 'same inp/state/head parameters and K evaluations; only baseline vs Fejer update differs',
139         'idea_grid': [{'cfg': q['cfg'], 'mean': q['result']['mean'],
140                        'per_seed': q['result']['per_seed']} for q in idea_candidates],
141         'mechanism_signature': mechanism_signature()})
142    Path('bench_report.json').write_text(json.dumps(report, indent=2))
143    print(json.dumps(report, indent=2))
144
145
146if __name__ == '__main__':
147    main()