Criticality-Guided Failure Replay / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10EPOCHS = 6
11BATCH = 128
12EPS = 0.02
13FAIL_THR = 1.5
14
15def seed_all(s):
16 random.seed(s); np.random.seed(s); torch.manual_seed(s)
17 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
18
19def proposal(c, alpha):
20 a = (EPS + np.clip(c, .01, .99)) ** alpha
21 z = float(a.mean())
22 q = a / a.sum()
23 w = z / a
24 ess = float((w.sum() ** 2) / (w @ w) / len(w))
25 return q, w, ess, z
26
27def math_check():
28 rng = np.random.default_rng(13)
29 c = rng.uniform(.01,.99,1000); y = rng.binomial(1,c)
30 g = rng.normal(size=(1000,7)); rows=[]
31 for alpha in (.5,1.,2.):
32 q,w,ess,z = proposal(c,alpha)
33 exact = (q[:,None]*w[:,None]*g).sum(0)
34 uniform = g.mean(0)
35 pred = (q*y).sum()/y.mean()
36 obs = float(y[rng.choice(len(y),20000,p=q)].mean()/y.mean())
37 rows.append({'alpha':alpha,'predicted_enrichment':float(pred),
38 'sampled_enrichment':obs,'identity_l2':float(np.linalg.norm(exact-uniform)),
39 'ess_fraction':ess})
40 return {'rows':rows,'max_identity_l2':max(x['identity_l2'] for x in rows),
41 'confirmed_identity':max(x['identity_l2'] for x in rows)<1e-12}
42
43def critic_scores(x,y,seed):
44 seed_all(seed+10000)
45 # auxiliary predictor sees state/window only; labels are eventual high-energy proxy
46 net=nn.Sequential(nn.Linear(x.shape[1],32),nn.Tanh(),nn.Linear(32,1))
47 opt=torch.optim.Adam(net.parameters(),lr=.01)
48 for _ in range(80):
49 loss=nn.functional.binary_cross_entropy_with_logits(net(x).squeeze(1),y)
50 opt.zero_grad(); loss.backward(); opt.step()
51 with torch.no_grad():
52 c=torch.sigmoid(net(x)).squeeze(1).numpy()
53 return np.clip(c,.01,.99), float(loss)
54
55def idea_train(seed,cfg,capture=False):
56 seed_all(seed)
57 d=get_dataset('dynamics',seed,n_train=400,n_test=200)
58 dev='cuda' if torch.cuda.is_available() else 'cpu'
59 # failure labels are held-out rollout outcome proxies from the training targets
60 labels=(d['ytr'].abs().max(dim=1).values.numpy()>FAIL_THR).astype(np.float32)
61 c,bce=critic_scores(d['xtr'],torch.tensor(labels),seed)
62 q,w,ess,z=proposal(c,cfg['alpha'])
63 try:
64 net=make_model('rnn_small',d['input_shape'],d['out_dim']).to(dev)
65 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
66 x,y=d['xtr'].to(dev),d['ytr'].to(dev)
67 for _ in range(EPOCHS):
68 net.train()
69 for _step in range((len(x)+BATCH-1)//BATCH):
70 ids=np.random.choice(len(x),BATCH,replace=True,p=q)
71 ix=torch.as_tensor(ids,device=dev)
72 per=(net(x[ix])-y[ix]).pow(2).mean(dim=1)
73 # finite-population p/q correction, self-normalized for stability
74 ww=torch.as_tensor(w[ids],dtype=per.dtype,device=dev)
75 loss=(per*ww).sum()/(ww.sum()+1e-8) if cfg['weighted'] else per.mean()
76 opt.zero_grad();loss.backward();opt.step()
77 net.eval()
78 with torch.no_grad(): metric=float((net(d['xte'].to(dev))-d['yte'].to(dev)).pow(2).mean())
79 except RuntimeError:
80 # explicit CPU fallback for tight/shared CUDA allocations
81 net=make_model('rnn_small',d['input_shape'],d['out_dim'])
82 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); x,y=d['xtr'],d['ytr']
83 for _ in range(EPOCHS):
84 for _step in range((len(x)+BATCH-1)//BATCH):
85 ids=np.random.choice(len(x),BATCH,replace=True,p=q); ix=torch.as_tensor(ids)
86 per=(net(x[ix])-y[ix]).pow(2).mean(1); ww=torch.tensor(w[ids],dtype=per.dtype)
87 loss=(per*ww).sum()/ww.sum() if cfg['weighted'] else per.mean()
88 opt.zero_grad();loss.backward();opt.step()
89 with torch.no_grad(): metric=float((net(d['xte'])-d['yte']).pow(2).mean())
90 if not capture:return metric
91 ids=np.random.choice(len(x),20000,replace=True,p=q)
92 observed=float(labels[ids].mean()/labels.mean())
93 # model-behaviour check: weighted and unweighted train-pool prediction losses
94 with torch.no_grad():
95 pred=(net(x).detach().cpu()-y.detach().cpu()).pow(2).mean(1).numpy()
96 weighted_mean=float((q*w*pred).sum()); uniform_mean=float(pred.mean())
97 return {'metric':metric,'critic_bce':bce,'failure_rate':float(labels.mean()),
98 'predicted_enrichment':float((q*labels).sum()/labels.mean()),
99 'observed_enrichment':observed,'ess_fraction':ess,
100 'weighted_pool_loss':weighted_mean,'uniform_pool_loss':uniform_mean}
101
102def baseline_train(seed,cfg):
103 seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=200)
104 _,metric,_=train_model(make_model('rnn_small',d['input_shape'],d['out_dim']),d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
105 return metric
106
107def main():
108 check=math_check()
109 grid=[{'lr':v,'epochs':EPOCHS,'alpha':0.,'weighted':False} for v in (.001,.003,.006)]
110 base=sweep_baseline(lambda cfg: (lambda s: baseline_train(s,cfg)),grid,seeds=SEEDS)
111 best_lr=base['best_cfg']['lr']
112 idea_cfgs=[{'lr':lr,'epochs':EPOCHS,'alpha':a,'weighted':True} for lr in (best_lr,.001,.006) for a in (.5,1.,2.)]
113 # keep a comparable 3-setting intervention sweep: best baseline lr, three replay strengths
114 idea_cfgs=[{'lr':best_lr,'epochs':EPOCHS,'alpha':a,'weighted':True} for a in (.5,1.,2.)]
115 idea_runs=[(c,evaluate(lambda s,c=c:idea_train(s,c),seeds=SEEDS)) for c in idea_cfgs]
116 cfg,idea=min(idea_runs,key=lambda z:z[1]['mean'])
117 sig=idea_train(0,cfg,True)
118 report=make_report('dynamics','rnn_small',base,idea,extra={
119 'mechanism_signature':sig,'math_check':check,
120 'prediction':'criticality proposal enrichment equals E_q[y]/E_p[y] while weighted expectation preserves uniform loss',
121 'confirmed': bool(check['confirmed_identity'] and abs(sig['observed_enrichment']-sig['predicted_enrichment'])/max(sig['predicted_enrichment'],1e-9)<.20),
122 'idea_sweep':[{'cfg':c,'result':r} for c,r in idea_runs],
123 'track_justification':'dynamics is the built-in structural match for rollout failure/control states'})
124 Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
125if __name__=='__main__': main()