Noisy Scrambling-Front Network / bench_experiment.py
Beats tuned baseline
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()