Pick-to-Learn Safety Fine-Tuning / bench_stage2.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7
  8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  9from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
 10
 11SEEDS = tuple(range(8))
 12TRACK = 'dynamics'
 13MODEL = 'rnn_small'
 14EPOCHS = 18
 15BATCH = 64
 16NTRAIN, NTEST = 1000, 300
 17# All learning rates are shared by baseline and idea, satisfying search parity.
 18GRID = [
 19    {'lr': 1e-3, 'weight_decay': 0.0},
 20    {'lr': 3e-3, 'weight_decay': 0.0},
 21    {'lr': 1e-2, 'weight_decay': 0.0},
 22]
 23# Safety constants fixed before running: large predicted angle is unsafe.
 24SAFETY_SCALE = 0.25
 25SAFETY_LIMIT = 0.80
 26BETA = 1.0
 27TEMP = 0.12
 28TOP_Q = 8
 29
 30
 31def seed_all(seed):
 32    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 33    if torch.cuda.is_available():
 34        try: torch.cuda.manual_seed_all(seed)
 35        except Exception: pass
 36    torch.set_num_threads(4)
 37
 38
 39def baseline_train(seed, cfg, keep=False):
 40    seed_all(seed)
 41    ds = get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST)
 42    net = make_model(MODEL, ds['input_shape'], ds['out_dim'])
 43    net, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'],
 44                                    batch=BATCH, weight_decay=cfg['weight_decay'],
 45                                    log=lambda *_: None)
 46    if keep: return metric, net, ds
 47    return metric
 48
 49
 50def pick_train(seed, cfg, keep=False):
 51    seed_all(seed)
 52    ds = get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST)
 53    net = make_model(MODEL, ds['input_shape'], ds['out_dim'])
 54    # The custom loop is the intervention: append worst violating trajectories
 55    # and optimize task MSE plus a differentiable surrogate safety penalty.
 56    try:
 57        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 58        net = net.to(device)
 59        x, y = ds['xtr'].to(device), ds['ytr'].to(device)
 60        opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 61        buffer = []
 62        for ep in range(EPOCHS):
 63            net.train(); perm = torch.randperm(len(x), device=device)
 64            for start in range(0, len(x), BATCH):
 65                idx = perm[start:start+BATCH]
 66                pred = net(x[idx]).reshape(-1)
 67                task = F.mse_loss(pred, y[idx].reshape(-1))
 68                # normalized violation v=[(|prediction|-limit)/scale]_+
 69                v = F.relu((pred.abs() - SAFETY_LIMIT) / SAFETY_SCALE)
 70                # Pick worst current trajectories, as prescribed by the idea.
 71                k = min(TOP_Q, len(idx))
 72                worst = torch.topk(v.detach(), k=k).indices
 73                buffer.extend(idx[worst].detach().cpu().tolist())
 74                if len(buffer) > 64: buffer = buffer[-64:]
 75                bx = x[torch.as_tensor(buffer, device=device)]
 76                bp = net(bx).reshape(-1)
 77                bv = F.relu((bp.abs() - SAFETY_LIMIT) / SAFETY_SCALE)
 78                safe = (F.softplus(bv / TEMP) * TEMP).pow(2).mean()
 79                loss = task + BETA * safe
 80                opt.zero_grad(); loss.backward(); opt.step()
 81        net.eval()
 82        with torch.no_grad():
 83            pred = net(ds['xte'].to(device)).reshape(-1)
 84            metric = float(F.mse_loss(pred, ds['yte'].to(device).reshape(-1)).cpu())
 85        if keep: return metric, net, ds
 86        return metric
 87    except RuntimeError:
 88        # Robust CPU fallback for a shared/fragile CUDA slot.
 89        seed_all(seed); torch.cuda.empty_cache() if torch.cuda.is_available() else None
 90        old = torch.cuda.is_available
 91        # Re-run identical intervention on CPU by temporarily selecting device locally.
 92        net = make_model(MODEL, ds['input_shape'], ds['out_dim'])
 93        x, y = ds['xtr'], ds['ytr']; opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'])
 94        buffer=[]
 95        for ep in range(EPOCHS):
 96            for start in range(0,len(x),BATCH):
 97                idx=torch.randperm(len(x))[start:start+BATCH]; pred=net(x[idx]).reshape(-1)
 98                v=F.relu((pred.abs()-SAFETY_LIMIT)/SAFETY_SCALE); k=min(TOP_Q,len(idx))
 99                buffer.extend(idx[torch.topk(v.detach(),k).indices].tolist()); buffer=buffer[-64:]
100                bp=net(x[torch.tensor(buffer)]).reshape(-1); bv=F.relu((bp.abs()-SAFETY_LIMIT)/SAFETY_SCALE)
101                loss=F.mse_loss(pred,y[idx].reshape(-1))+BETA*(F.softplus(bv/TEMP)*TEMP).pow(2).mean()
102                opt.zero_grad(); loss.backward(); opt.step()
103        with torch.no_grad(): metric=float(F.mse_loss(net(ds['xte']).reshape(-1),ds['yte'].reshape(-1)))
104        return (metric,net,ds) if keep else metric
105
106
107def behavior_signature(base_cfg, idea_cfg):
108    rows=[]
109    for s in SEEDS:
110        bm,bn,bd=baseline_train(s,base_cfg,True); im,inn,idd=pick_train(s,idea_cfg,True)
111        with torch.no_grad():
112            bp=bn(bd['xte'].to(next(bn.parameters()).device)).reshape(-1).cpu().numpy()
113            ip=inn(idd['xte'].to(next(inn.parameters()).device)).reshape(-1).cpu().numpy()
114        true=bd['yte'].numpy().reshape(-1)
115        rows.append({'seed':s,'baseline_pred_rate':float(np.mean(np.abs(bp)>SAFETY_LIMIT)),
116                     'idea_pred_rate':float(np.mean(np.abs(ip)>SAFETY_LIMIT)),
117                     'observed_rate':float(np.mean(np.abs(true)>SAFETY_LIMIT)),
118                     'baseline_pred_max_margin':float(np.max(np.maximum(np.abs(bp)-SAFETY_LIMIT,0))),
119                     'idea_pred_max_margin':float(np.max(np.maximum(np.abs(ip)-SAFETY_LIMIT,0)))})
120    obs=float(np.mean([r['observed_rate'] for r in rows]))
121    pred_b=float(np.mean([r['baseline_pred_rate'] for r in rows])); pred_i=float(np.mean([r['idea_pred_rate'] for r in rows]))
122    # Prediction tested at NN scale: adaptive training should reduce predicted and
123    # observed rare-event rates; confirmed only if both decrease by >=20%.
124    return {'safety_limit':SAFETY_LIMIT,'scale':SAFETY_SCALE,'rows':rows,
125            'observed_rate_mean':obs,'predicted_rate_baseline_mean':pred_b,
126            'predicted_rate_idea_mean':pred_i,
127            'predicted_reduction_fraction':float((pred_b-pred_i)/max(pred_b,1e-9)),
128            'observed_reduction_fraction': 0.0,
129            'confirmed': False}
130
131
132def main():
133    base = sweep_baseline(lambda cfg: (lambda seed: baseline_train(seed,cfg)), GRID, seeds=(0,1,2,3))
134    # Three idea settings: baseline-best plus two nearby settings, all in GRID.
135    idea_cfgs = GRID
136    idea_runs=[]
137    for cfg in idea_cfgs:
138        r={'cfg':cfg,'result':{'mean':0,'std':0,'per_seed':[],'n':0}}
139        vals=[pick_train(s,cfg) for s in SEEDS]
140        r['result']={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':[float(v) for v in vals],'n':len(vals)}
141        idea_runs.append(r)
142    best=min(idea_runs,key=lambda z:z['result']['mean'])
143    sig=behavior_signature(base['best_cfg'],best['cfg'])
144    report=make_report(TRACK,MODEL,base,best['result'],sig)
145    report['idea_sweep']=idea_runs
146    report['protocol']={'seeds':list(SEEDS),'n_train':NTRAIN,'n_test':NTEST,'epochs':EPOCHS,'batch':BATCH,
147                        'structural_match':'dynamics: controlled pendulum rollout and stability/safety violations',
148                        'baseline_method':'uniform per-example MSE','idea_method':'top-8 normalized predicted-angle violation replay'}
149    Path('bench_report.json').write_text(json.dumps(report,indent=2))
150    print(json.dumps(report,indent=2))
151
152if __name__=='__main__': main()