Residual-Scenario Safety Training / multi_seed.py
Failed on benchmark
1import json, random
2import numpy as np
3import torch
4from experiment import make_data, fit_nominal, fit_scenario, predict, violation_metrics, exact_slack_check, device
5
6rows=[]
7for seed in [535, 536, 537, 538, 539]:
8 np.random.seed(seed); random.seed(seed); torch.manual_seed(seed)
9 rng=np.random.default_rng(seed)
10 xtr,ytr,cleantr=make_data(700,.10,rng)
11 xte,yte,cleante=make_data(5000,.10,rng)
12 wb=fit_nominal(xtr,ytr)
13 rb=ytr-predict(wb,xtr)
14 ws=fit_scenario(xtr,ytr,rb)
15 b=violation_metrics(predict(wb,xte),cleante,.10,np.random.default_rng(seed+2))
16 s=violation_metrics(predict(ws,xte),cleante,.10,np.random.default_rng(seed+2))
17 rows.append({'seed':seed,'baseline':b,'scenario':s})
18
19def avg(key, method):
20 return float(np.mean([r[method][key] for r in rows]))
21summary={}
22for key in ['nominal_mse_to_clean','p95_abs_error','trajectory_sample_violation_rate','mean_slack']:
23 summary[key]={'baseline_mean':avg(key,'baseline'),'scenario_mean':avg(key,'scenario'),
24 'relative_change':avg(key,'scenario')/avg(key,'baseline')-1}
25print(json.dumps({'device':str(device),'slack_check':exact_slack_check(),'runs':rows,'summary':summary},indent=2))