Transient-risk certificate for Langevin training / local_nn_bench.py
Mechanism confirmed, baseline not beaten
1import json, math, itertools
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEEDS = list(range(8))
8# Shared search space: every idea setting is evaluated by baseline too.
9CONFIGS = [(lr, noise) for lr in (0.01, 0.03, 0.05) for noise in (0.005, 0.02)]
10EPOCHS, BATCH = 24, 64
11
12class MLP(nn.Module):
13 def __init__(self):
14 super().__init__()
15 self.net = nn.Sequential(nn.Linear(10, 32), nn.Tanh(), nn.Linear(32, 1))
16 def forward(self, x): return self.net(x)
17
18def data(seed, n=400):
19 r = np.random.default_rng(seed)
20 x = r.normal(size=(n, 10)).astype('float32')
21 y = (np.sin(x[:,0]) + .5*x[:,1]**2 - .3*x[:,2] + .15*x[:,3]*x[:,4] + .1*r.normal(size=n)).astype('float32')
22 ix = r.permutation(n); tr, te = ix[:300], ix[300:]
23 return torch.tensor(x[tr]), torch.tensor(y[tr,None]), torch.tensor(x[te]), torch.tensor(y[te,None])
24
25def risk_stats(model, center, threshold, pi_floor=1e-3):
26 # A is a deliberately predefined unsafe parameter region: distance from
27 # the initial basin exceeds threshold. Empirical stationary risk is the
28 # tail of a local Gaussian fit to parameter perturbation samples.
29 with torch.no_grad():
30 d = torch.cat([(p-center[i]).flatten() for i,p in enumerate(model.parameters())])
31 radius = float(torch.linalg.vector_norm(d))
32 # local quadratic proxy: stationary radius scale estimated from current weights
33 scale = max(float(torch.linalg.vector_norm(torch.cat([p.flatten() for p in model.parameters()])))/12., 1e-3)
34 z = max((threshold-radius)/scale, -8.)
35 pi = float(.5*math.erfc(z/math.sqrt(2)))
36 return min(max(pi, pi_floor), 1.-pi_floor), radius
37
38def train(seed, lr, noise, controlled):
39 torch.manual_seed(seed); np.random.seed(seed)
40 x,y,xe,ye = data(seed)
41 model=MLP(); center=[p.detach().clone() for p in model.parameters()]
42 # Threshold makes excursions measurable but not trivially impossible.
43 threshold=1.8
44 # Conservative certificate parameters (m is discrete relaxation proxy).
45 m=.12; delta=.18; chi2=9.0; stop=None; unsafe=[]
46 gen=torch.Generator().manual_seed(seed+1000)
47 for ep in range(EPOCHS):
48 perm=torch.randperm(len(x), generator=gen)
49 for st in range(0,len(x),BATCH):
50 pred=model(x[perm[st:st+BATCH]]); loss=((pred-y[perm[st:st+BATCH]])**2).mean()
51 model.zero_grad(); loss.backward()
52 use_noise = noise if (not controlled or stop is None) else 0.
53 with torch.no_grad():
54 for p in model.parameters():
55 if p.grad is not None:
56 p -= lr*p.grad + math.sqrt(2*lr*use_noise)*torch.randn_like(p)
57 pi, radius=risk_stats(model, center, threshold)
58 bound=min(1., pi+math.sqrt(pi*chi2)*math.exp(-m*(ep+1)))
59 unsafe.append(float(radius>threshold))
60 if controlled and stop is None and bound <= delta: stop=ep+1
61 with torch.no_grad():
62 mse=float(((model(xe)-ye)**2).mean())
63 return {'mse':mse,'max_unsafe':max(unsafe),'mean_unsafe':float(np.mean(unsafe)),
64 'post_stop_unsafe':float(np.mean(unsafe[stop-1:])) if stop else float(np.mean(unsafe)),
65 'stop_epoch':stop,'final_radius':risk_stats(model,center,threshold)[1],
66 'predicted_bound_final':bound}
67
68def mean_metric(rows,k): return float(np.mean([r[k] for r in rows]))
69def perm_p(a,b):
70 d=np.array([b[i]-a[i] for i in range(len(a))]); obs=abs(d.mean()); rng=np.random.default_rng(991)
71 cnt=0; n=20000
72 for _ in range(n):
73 if abs((d*rng.choice([-1,1],len(d))).mean())>=obs: cnt+=1
74 return float((cnt+1)/(n+1))
75
76def main():
77 allres={}
78 for cfg in CONFIGS:
79 lr,noise=cfg
80 base=[train(s,lr,noise,False) for s in SEEDS]
81 idea=[train(s,lr,noise,True) for s in SEEDS]
82 allres[f'{lr}_{noise}']={'lr':lr,'noise':noise,'baseline':base,'idea':idea,
83 'baseline_mse':mean_metric(base,'mse'),'idea_mse':mean_metric(idea,'mse')}
84 # baseline sweep chooses lowest mean MSE; idea reports its best setting,
85 # both on the same union of configurations.
86 best_base=min(allres.values(),key=lambda z:z['baseline_mse'])
87 best_idea=min(allres.values(),key=lambda z:z['idea_mse'])
88 a=best_base['baseline']; b=best_idea['idea']
89 delta=mean_metric(b,'mse')-mean_metric(a,'mse')
90 # Signature is behavior measured on trained models, not an identity.
91 pred=float(np.mean([r['predicted_bound_final'] for r in b]))
92 obs=float(np.mean([r['mean_unsafe'] for r in b]))
93 sig={'predicted_final_bound':pred,'observed_mean_unsafe_frequency':obs,
94 'ratio_observed_to_bound':obs/max(pred,1e-9),
95 'confirmed': bool(obs <= pred*1.25 + .02)}
96 report={'track':'local_friedman_mlp_optimizer','task':'regression','metric':'test_mse',
97 'baseline_sweep':{k:{'mean_mse':v['baseline_mse'],'lr':v['lr'],'noise':v['noise']} for k,v in allres.items()},
98 'best_baseline':{'lr':best_base['lr'],'noise':best_base['noise'],'per_seed':a},
99 'idea_best':{'lr':best_idea['lr'],'noise':best_idea['noise'],'per_seed':b},
100 'paired_delta_mean_idea_minus_baseline':delta,'permutation_p_value':perm_p([r['mse'] for r in a],[r['mse'] for r in b]),
101 'mechanism_signature':sig,'all_configs':allres,
102 'custom_track':{'name':'local_friedman_mlp_optimizer','file':'local_nn_bench.py','domain':'tabular/optimizer'},
103 'note':'Official bench path /home/maxwelhelp/all/math2nn/bench and README were absent; this is a fallback, not official bench evidence.'}
104 Path('bench_report.json').write_text(json.dumps(report,indent=2))
105 print(json.dumps({'delta':delta,'p':report['permutation_p_value'],'signature':sig,'best_base':best_base['baseline_mse'],'best_idea':best_idea['idea_mse']},indent=2))
106if __name__=='__main__': main()