Criticality-Guided Failure Replay / stage2_bench.py

Mechanism confirmed, baseline not beaten

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