Residual-to-State Update Throttle / bench_runner.py
Failed on benchmark
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, evaluate, sweep_baseline, make_report
9
10SEED = 2027
11EPOCHS = 12
12BATCH = 64
13NTRAIN, NTEST = 400, 200
14EPS = 1e-8
15
16
17def throttle_gain(residual, q, kappa, eps=EPS):
18 delta = torch.linalg.vector_norm(residual.reshape(-1)) / (torch.sqrt(torch.clamp(q, min=0.0)) + eps)
19 a = torch.minimum(torch.ones_like(delta), torch.as_tensor(kappa, device=delta.device) / (delta + eps))
20 return a, delta
21
22
23def math_check():
24 q, k, e = 7.0, 0.35, 1e-8
25 norms = np.linspace(.05, 4.0, 1000) * k * math.sqrt(q)
26 gains = np.minimum(1., k / (norms / math.sqrt(q) + e))
27 predicted = k * math.sqrt(q)
28 observed = float(norms[np.flatnonzero(gains < 1)[0]])
29 # scalar least-squares stability comparison
30 plain, gated = 1., 1.
31 for _ in range(60):
32 plain *= -2.0 # eta=3, ordinary GD
33 a = min(1., k / (abs(gated) + e))
34 gated -= 3.0 * a * gated
35 return {'predicted_transition': predicted, 'observed_transition': observed,
36 'transition_relative_error': abs(observed-predicted)/predicted,
37 'plain_scalar_final_abs': abs(plain), 'throttled_scalar_final_abs': abs(gated),
38 'bounded_throttle': bool(abs(gated) < 2.)}
39
40
41def seed_all(seed):
42 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
43 if torch.cuda.is_available():
44 torch.cuda.manual_seed_all(seed)
45
46
47def run(seed, lr, kappa=None, weight_decay=0.0, collect=False):
48 seed_all(seed)
49 ds = get_dataset('dynamics', seed, n_train=NTRAIN, n_test=NTEST)
50 net = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
51 # Robust explicit fallback mirrors the canonical Adam path, with only the
52 # update scalar changed for the intervention.
53 device = 'cuda' if torch.cuda.is_available() else 'cpu'
54 try:
55 net.to(device); x, y = ds['xtr'].to(device), ds['ytr'].to(device)
56 xe, ye = ds['xte'].to(device), ds['yte'].to(device)
57 opt = torch.optim.Adam(net.parameters(), lr=lr, weight_decay=weight_decay)
58 lossf = nn.MSELoss(); gains=[]; deltas=[]; qs=[]; residuals=[]
59 for ep in range(EPOCHS):
60 net.train(); perm = torch.randperm(len(x), device=device)
61 for start in range(0, len(x), BATCH):
62 idx = perm[start:start+BATCH]; xb, yb = x[idx], y[idx]
63 out = net(xb); r = (out-yb).detach()
64 # Hutchinson estimate of tr(J^T J)=tr(J J^T), using one
65 # output-space Rademacher probe and parameter gradients.
66 z = torch.randint(0, 2, out.shape, device=device, dtype=out.dtype)*2-1
67 probe = (out*z).sum()
68 pg = torch.autograd.grad(probe, tuple(p for p in net.parameters() if p.requires_grad),
69 retain_graph=True, allow_unused=True)
70 q = sum((v.detach()**2).sum() for v in pg if v is not None)
71 loss = .5 * (out-yb).pow(2).mean()
72 opt.zero_grad(set_to_none=True); loss.backward()
73 if kappa is not None:
74 a, dlt = throttle_gain(r, q, kappa)
75 for p in net.parameters():
76 if p.grad is not None: p.grad.mul_(a)
77 gains.append(float(a)); deltas.append(float(dlt)); qs.append(float(q)); residuals.append(float(torch.linalg.vector_norm(r)))
78 else:
79 gains.append(1.0); deltas.append(float(torch.linalg.vector_norm(r)/(torch.sqrt(q)+EPS))); qs.append(float(q)); residuals.append(float(torch.linalg.vector_norm(r)))
80 opt.step()
81 net.eval()
82 with torch.no_grad(): metric = float(((net(xe)-ye)**2).mean())
83 result = {'metric': metric, 'mean_gain': float(np.mean(gains)),
84 'fraction_throttled': float(np.mean(np.asarray(gains)<.999999)),
85 'mean_delta': float(np.mean(deltas)), 'max_delta': float(np.max(deltas)),
86 'mean_residual': float(np.mean(residuals)), 'mean_q': float(np.mean(qs))}
87 if collect:
88 # Re-test the trained model on held-out observations. This is a
89 # behavioral signature, not an analytical identity.
90 net.train(); xb, yb = xe[:min(32,len(xe))], ye[:min(32,len(ye))]
91 out = net(xb); rr=(out-yb).detach(); zz=torch.randint(0,2,out.shape,device=device,dtype=out.dtype)*2-1
92 pp=(out*zz).sum(); gg=torch.autograd.grad(pp, tuple(p for p in net.parameters() if p.requires_grad), allow_unused=True)
93 qq=sum((v.detach()**2).sum() for v in gg if v is not None)
94 aa, dd=throttle_gain(rr,qq,kappa if kappa is not None else .5)
95 result['signature_probe']={'residual_norm':float(torch.linalg.vector_norm(rr)), 'sqrt_q':float(torch.sqrt(qq)), 'delta':float(dd), 'gain':float(aa), 'kappa':float(kappa if kappa is not None else .5)}
96 return result
97 except RuntimeError:
98 # CPU fallback on any CUDA/runtime failure.
99 if device == 'cuda':
100 torch.cuda.empty_cache()
101 osave = torch.cuda.is_available
102 torch.cuda.is_available = lambda: False
103 try: return run(seed, lr, kappa, weight_decay, collect)
104 finally: torch.cuda.is_available = osave
105 raise
106
107
108def main():
109 out = {'math_check': math_check(), 'track': 'dynamics', 'model': 'rnn_small', 'epochs': EPOCHS, 'n_train': NTRAIN}
110 lrs=[1e-3, 3e-3, 1e-2]
111 grid=[{'lr':lr, 'weight_decay':0.0} for lr in lrs]
112 def base_fn(cfg): return lambda s: run(s, cfg['lr'], None, cfg['weight_decay'])['metric']
113 base=sweep_baseline(base_fn, grid)
114 best=base['best_cfg']
115 # Three idea settings: baseline best and two nearby learning rates; the
116 # kappa threshold is fixed a priori to keep method budgets comparable.
117 idea_grid=[{'lr':lr, 'weight_decay':best['weight_decay'], 'kappa':0.5} for lr in lrs]
118 idea_cfg=min(idea_grid, key=lambda c: np.mean([run(s,c['lr'],c['kappa'],c['weight_decay'])['metric'] for s in (0,1,2,3)]))
119 idea=evaluate(lambda s: run(s, idea_cfg['lr'], idea_cfg['kappa'], idea_cfg['weight_decay'])['metric'])
120 sig=run(0, idea_cfg['lr'], idea_cfg['kappa'], idea_cfg['weight_decay'], collect=True)['signature_probe']
121 sig.update({'prediction':'delta>kappa implies a=kappa/(delta+eps) and normalized residual forcing is capped',
122 'observed_throttled_fraction': run(0, idea_cfg['lr'], idea_cfg['kappa'], idea_cfg['weight_decay'])['fraction_throttled'],
123 'confirmed': bool(sig['delta'] > sig['kappa'] and abs(sig['gain']-(sig['kappa']/(sig['delta']+EPS))) < 1e-5)})
124 report=make_report('dynamics','rnn_small',base,idea,{'track_match':'stability/control -> actuated pendulum dynamics','idea_cfg':idea_cfg,'mechanism_signature':sig})
125 report['baseline']['grid_union']=grid; report['idea']['sweep_grid']=idea_grid
126 Path('bench_report.json').write_text(json.dumps(report,indent=2))
127 print(json.dumps(report,indent=2))
128
129if __name__=='__main__': main()