Transient-risk certificate for Langevin training / local_nn_bench.py

Mechanism confirmed, baseline not beaten

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