Certified Temporal Budget for Neural Control / bench_experiment.py
Mechanism confirmed, baseline not beaten
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()