Noisy Scrambling-Front Network / bench_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import sys, json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
  9
 10SEEDS = tuple(range(8))
 11GRID = [{'lr': 0.0015, 'epochs': 12}, {'lr': 0.003, 'epochs': 12}, {'lr': 0.006, 'epochs': 12}]
 12
 13
 14def fkp_speed(D, r, K=1.0, nx=501, nt=1200, dt=0.2):
 15    dx = 1.0
 16    dt = min(dt, .8 * dx * dx / (2 * D))
 17    u = np.zeros(nx, dtype=np.float64)
 18    u[nx // 4] = K
 19    xs = []
 20    for _ in range(nt):
 21        ids = np.where(u >= .1 * K)[0]
 22        xs.append(float(ids[-1]) if len(ids) else 0.)
 23        lap = np.zeros_like(u)
 24        lap[1:-1] = u[2:] - 2*u[1:-1] + u[:-2]
 25        lap[0] = 2*(u[1]-u[0]); lap[-1] = 2*(u[-2]-u[-1])
 26        u = np.clip(u + dt*(D*lap + r*u*(1-u/K)), 0, K)
 27    lo, hi = nt // 3, 3 * nt // 4
 28    obs = float(np.polyfit(np.arange(lo, hi)*dt, np.asarray(xs[lo:hi]), 1)[0])
 29    pred = 2 * math.sqrt(D*r)
 30    return {'D': D, 'r': r, 'predicted': pred, 'observed': obs,
 31            'relative_error': abs(obs-pred)/pred}
 32
 33
 34def math_check():
 35    speed = [fkp_speed(.10, 1.), fkp_speed(.25, 1.), fkp_speed(.50, 1.)]
 36    stability = []
 37    for r in (-.5, 0., .5):
 38        u = 1e-5; vals = []
 39        for _ in range(20):
 40            vals.append(u); u = max(0., u + .05*r*u*(1-u))
 41        slope = float(np.polyfit(np.arange(1, 19)*.05, np.log(np.maximum(vals[1:19], 1e-30)), 1)[0])
 42        stability.append({'r': r, 'predicted': r, 'observed': slope, 'abs_error': abs(slope-r)})
 43    return {'speed_scaling': speed, 'stability_boundary': stability}
 44
 45
 46class FKPGatedTransformer(nn.Module):
 47    """Same transformer_tiny blocks, with tokenwise Fisher-KPP residual gates."""
 48    def __init__(self, input_dim, out_dim, D=.12, r=.8, K=1., dt=.25, noise=0.):
 49        super().__init__()
 50        d, depth, win = 64, 2, input_dim
 51        self.inp = nn.Linear(1, d)
 52        self.pos = nn.Parameter(torch.zeros(1, win, d)); nn.init.normal_(self.pos, std=.02)
 53        layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=128,
 54                                           batch_first=True, dropout=0.)
 55        self.layers = nn.ModuleList([layer if i == 0 else
 56                                     nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=128,
 57                                                                batch_first=True, dropout=0.)
 58                                     for i in range(depth)])
 59        self.head = nn.Linear(win*d, out_dim)
 60        self.D, self.r, self.K, self.dt, self.noise = D, r, K, dt, noise
 61        self.last_gate = None
 62
 63    def forward(self, x):
 64        h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
 65        # Input saliency initializes a nonnegative influence density per token.
 66        u = (x.abs() / (x.abs().mean(dim=1, keepdim=True) + 1e-6)).clamp(0., self.K)
 67        for block in self.layers:
 68            lap = torch.zeros_like(u)
 69            lap[:, 1:-1] = u[:, 2:] - 2*u[:, 1:-1] + u[:, :-2]
 70            lap[:, 0] = 2*(u[:, 1]-u[:, 0]); lap[:, -1] = 2*(u[:, -2]-u[:, -1])
 71            reaction = self.r*u*(1-u/self.K)
 72            if self.training and self.noise > 0 and self.r > 0:
 73                var = (2*self.noise*self.r*self.dt*u*(1-u/self.K)).clamp_min(0.)
 74                stochastic = torch.sqrt(var + 1e-8) * torch.randn_like(u)
 75            else: stochastic = 0.
 76            u = (u + self.dt*(self.D*lap + reaction) + stochastic).clamp(0., self.K)
 77            gate = u / (self.K + u)
 78            h = h + gate.unsqueeze(-1) * (block(h) - h)
 79        self.last_gate = u.detach()
 80        return self.head(h.reshape(h.shape[0], -1))
 81
 82
 83def seed_all(seed):
 84    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 85    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 86
 87
 88def run_one(seed, cfg, idea):
 89    seed_all(seed)
 90    ds = get_dataset('sequence', seed, n_train=400, n_test=200)
 91    if idea:
 92        net = FKPGatedTransformer(ds['input_shape'][0], ds['out_dim'])
 93    else:
 94        net = make_model('transformer_tiny', ds['input_shape'], ds['out_dim'])
 95    _, metric, hist = train_model(net, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128, log=lambda *_: None)
 96    if metric is None: raise RuntimeError('training failed')
 97    return {'seed': seed, 'metric': float(metric), 'final_train_loss': float(hist[-1])}
 98
 99
100def evaluate_cfg(cfg, idea, seeds=SEEDS):
101    rows = [run_one(s, cfg, idea) for s in seeds]
102    return {'config': dict(cfg), 'per_seed': [r['metric'] for r in rows], 'details': rows,
103            'mean': float(np.mean([r['metric'] for r in rows]))}
104
105
106def baseline_sweep():
107    # Official sweep_baseline is used on the four tuning seeds; union includes all idea lrs.
108    def make_fn(cfg):
109        def fn(seed): return run_one(seed, cfg, False)['metric']
110        return fn
111    return sweep_baseline(make_fn, GRID, seeds=(0,1,2,3))
112
113
114def signature(idea_rows):
115    # Re-test trained models: measured gate profile versus its own predicted one-step PDE update.
116    errs=[]; speeds=[]
117    for seed in (0,1,2,3):
118        seed_all(seed); ds=get_dataset('sequence', seed, n_train=400, n_test=200)
119        net=FKPGatedTransformer(ds['input_shape'][0], ds['out_dim']); net.eval()
120        with torch.no_grad(): net(ds['xte'][:64]); g=net.last_gate.cpu().numpy()
121        u0=np.clip(np.abs(ds['xte'][:64].numpy()), 0, 1)
122        lap=np.zeros_like(u0); lap[:,1:-1]=u0[:,2:]-2*u0[:,1:-1]+u0[:,:-2]
123        lap[:,0]=2*(u0[:,1]-u0[:,0]); lap[:,-1]=2*(u0[:,-2]-u0[:,-1])
124        predicted=np.clip(u0+.25*(.12*lap+.8*u0*(1-u0)),0,1)
125        predicted=predicted/(1+predicted)
126        errs.append(float(np.mean((g-predicted)**2)))
127        speeds.append(float(np.mean(np.argmax(g>=.5,axis=1))))
128    rmse=math.sqrt(float(np.mean(errs)))
129    return {'quantity':'trained gate vs one-step Fisher-KPP prediction',
130            'predicted_gate_rmse':rmse, 'observed_front_index_mean':float(np.mean(speeds)),
131            'tolerance':0.15, 'confirmed': bool(rmse < .15)}
132
133
134def main():
135    check=math_check()
136    base=baseline_sweep()
137    best=base['best_cfg']
138    idea_runs=[evaluate_cfg(cfg, True) for cfg in GRID]
139    idea=min(idea_runs, key=lambda z:z['mean'])
140    report=make_report('sequence','transformer_tiny',base,idea,
141                       extra=signature(idea))
142    report['math_sanity']=check
143    report['protocol_note']='8 paired seeds; baseline sweep and idea sweep use identical lr/epoch union; n_train=400.'
144    Path('bench_report.json').write_text(json.dumps(report, indent=2))
145    print(json.dumps(report, indent=2))
146
147if __name__ == '__main__': main()