Residual-Scenario Safety Training / multi_seed.py

Failed on benchmark

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