Certified Temporal Budget for Neural Control / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, sys, random
  2import numpy as np
  3import torch
  4import torch.nn.functional as F
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7import bench
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = tuple(range(4))
 11EPOCHS = 12
 12BATCH = 128
 13THETA_LIMIT = 1.5
 14HORIZON = 0.20
 15
 16
 17def seed_all(seed):
 18    random.seed(seed)
 19    np.random.seed(seed)
 20    torch.manual_seed(seed)
 21    if torch.cuda.is_available():
 22        try:
 23            torch.cuda.manual_seed_all(seed)
 24        except Exception:
 25            pass
 26
 27
 28def baseline_train(cfg, seed):
 29    seed_all(seed)
 30    ds = bench.get_dataset('dynamics', seed, 400, 100)
 31    model = bench.make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 32    _, metric, _ = bench.train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
 33                                     batch=BATCH, weight_decay=cfg['weight_decay'],
 34                                     log=lambda *a, **k: None)
 35    return float(metric)
 36
 37
 38def certificate_terms(x, pred):
 39    last = x[:, -3:]
 40    theta, omega, action = last[:, 0], last[:, 1], last[:, 2]
 41    phi0 = THETA_LIMIT - torch.abs(theta)
 42    gamma = torch.abs(omega) + 0.20 * (torch.abs(action) + 1.0) + 0.50
 43    phi_pred = THETA_LIMIT - torch.abs(pred[:, 0])
 44    contract = phi0 - gamma * HORIZON
 45    return phi0, gamma, phi_pred, contract
 46
 47
 48def idea_train(cfg, seed, return_signature=False):
 49    seed_all(seed)
 50    ds = bench.get_dataset('dynamics', seed, 400, 100)
 51    model = bench.make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 52    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 53    try:
 54        model.to(device)
 55        xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
 56        xte, yte = ds['xte'].to(device), ds['yte'].to(device)
 57        opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 58        model.train()
 59        for _ in range(EPOCHS):
 60            perm = torch.randperm(len(xtr), device=device)
 61            for j in range(0, len(xtr), BATCH):
 62                xb, yb = xtr[perm[j:j+BATCH]], ytr[perm[j:j+BATCH]]
 63                pred = model(xb)
 64                phi0, gamma, phi_pred, contract = certificate_terms(xb, pred)
 65                safety_penalty = F.relu(-contract).pow(2).mean()
 66                loss = F.mse_loss(pred, yb) + cfg['cert_weight'] * safety_penalty
 67                opt.zero_grad(); loss.backward(); opt.step()
 68        model.eval()
 69        with torch.no_grad():
 70            pred = model(xte)
 71            mse = F.mse_loss(pred, yte).item()
 72            phi0, gamma, phi_pred, contract = certificate_terms(xte, pred)
 73            predicted_drop = (gamma * HORIZON).cpu().numpy()
 74            observed_drop = (phi0 - phi_pred).cpu().numpy()
 75            residual = observed_drop - predicted_drop
 76            permitted = (contract >= 0).cpu().numpy()
 77            violations = ((phi_pred < 0) & permitted).sum()
 78        result = float(mse)
 79        if return_signature:
 80            return result, {
 81                'n_test': int(len(xte)),
 82                'predicted_certificate_drop_mean': float(np.mean(predicted_drop)),
 83                'observed_certificate_drop_mean': float(np.mean(observed_drop)),
 84                'predicted_vs_observed_drop_ratio': float(np.mean(observed_drop) / (np.mean(predicted_drop) + 1e-12)),
 85                'contract_permitted_fraction': float(np.mean(permitted)),
 86                'permitted_certificate_violations': int(violations),
 87                'confirmed': bool(np.mean(observed_drop) <= np.mean(predicted_drop) * 1.20 + 1e-8)
 88            }
 89        return result
 90    except Exception:
 91        if device.type == 'cuda':
 92            torch.cuda.empty_cache()
 93        # Re-run entirely on CPU after any CUDA/runtime failure.
 94        torch.set_default_device('cpu')
 95        return idea_train_cpu(cfg, seed, return_signature)
 96
 97
 98def idea_train_cpu(cfg, seed, return_signature=False):
 99    seed_all(seed)
100    ds = bench.get_dataset('dynamics', seed, 400, 100)
101    model = bench.make_model('rnn_small', ds['input_shape'], ds['out_dim'])
102    xtr, ytr, xte, yte = ds['xtr'], ds['ytr'], ds['xte'], ds['yte']
103    opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
104    for _ in range(EPOCHS):
105        for j in range(0, len(xtr), BATCH):
106            pred = model(xtr[j:j+BATCH]); _, _, _, c = certificate_terms(xtr[j:j+BATCH], pred)
107            loss = F.mse_loss(pred, ytr[j:j+BATCH]) + cfg['cert_weight'] * F.relu(-c).pow(2).mean()
108            opt.zero_grad(); loss.backward(); opt.step()
109    with torch.no_grad():
110        pred = model(xte); mse = F.mse_loss(pred, yte).item()
111        phi0, gamma, phi_pred, contract = certificate_terms(xte, pred)
112        pd = (gamma*HORIZON).numpy(); od = (phi0-phi_pred).numpy(); permitted = (contract >= 0).numpy()
113    if return_signature:
114        return mse, {'n_test': len(xte), 'predicted_certificate_drop_mean': float(pd.mean()), 'observed_certificate_drop_mean': float(od.mean()), 'predicted_vs_observed_drop_ratio': float(od.mean()/(pd.mean()+1e-12)), 'contract_permitted_fraction': float(permitted.mean()), 'permitted_certificate_violations': int(((phi_pred < 0).numpy() & permitted).sum()), 'confirmed': bool(od.mean() <= pd.mean()*1.20+1e-8)}
115    return mse
116
117
118def main():
119    # Union of all idea learning rates is included in the baseline sweep.
120    grid = [{'lr': lr, 'weight_decay': wd} for lr in (0.0015, 0.003, 0.006) for wd in (0.0, 1e-4)]
121    base = bench.sweep_baseline(lambda c: lambda s: baseline_train(c, s), grid, seeds=SWEEP_SEEDS)
122    idea_cfgs = [dict(base['best_cfg'], cert_weight=w) for w in (0.01, 0.05, 0.20)]
123    idea_trials = []
124    for c in idea_cfgs:
125        vals = [idea_train(c, s) for s in SWEEP_SEEDS]
126        idea_trials.append({'cfg': c, 'mean': float(np.mean(vals))})
127    best = min(idea_trials, key=lambda z: z['mean'])
128    best_cfg = best['cfg']
129    vals = [idea_train(best_cfg, s) for s in SEEDS]
130    idea = {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals), 'best_cfg': best_cfg, 'sweep': idea_trials}
131    signatures = [idea_train(best_cfg, s, True)[1] for s in SEEDS]
132    sig = {k: (float(np.mean([x[k] for x in signatures])) if isinstance(signatures[0][k], (float, int)) and k not in ('n_test','permitted_certificate_violations') else (int(np.sum([x[k] for x in signatures])) if k in ('n_test','permitted_certificate_violations') else bool(all(x[k] for x in signatures)))) for k in signatures[0]}
133    report = bench.make_report('dynamics', 'rnn_small', base, idea, {'predicted_vs_observed': sig, 'confirmed': sig['confirmed']})
134    with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2)
135    print(json.dumps(report, indent=2))
136
137if __name__ == '__main__':
138    main()