Parabolic Riesz Feature Preconditioner / stage2_bench.py

Mechanism confirmed, baseline not beaten

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