Download scripts/run_v10_audit.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 3.49 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v10_audit.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_v10_audit.py
-
curl -L -o run_v10_audit.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v10_audit.py
3.49 kB
| """V10: frozen-model counterfactual benchmark integrity and physics audit.""" | |
| from __future__ import annotations | |
| import argparse,csv,json | |
| from pathlib import Path | |
| from statistics import mean,stdev | |
| import torch | |
| from spectral_world_models.models import build_model | |
| from spectral_world_models.train import train_one_model | |
| from spectral_world_models.v9_causal_generalization import generate_family_npz | |
| from spectral_world_models.v10_cf_audit import REGIMES,build_oracle_audit,assert_audit_integrity,evaluate_counterfactual_audit | |
| DEFAULT_MODELS=['neural_operator','swm_selective','swm_structured_cf','swm_structured_cf_v9','no_spectral_transition'] | |
| def stats(x): return mean(x),stdev(x) if len(x)>1 else 0. | |
| def main(): | |
| ap=argparse.ArgumentParser(description='V10 counterfactual integrity audit with regime-specific oracle trajectory hashes.') | |
| ap.add_argument('--epochs',type=int,default=3); ap.add_argument('--batch-size',type=int,default=64); ap.add_argument('--seeds',type=int,nargs='+',default=[0,1,2,3,4]); ap.add_argument('--cf-horizon',type=int,default=30); ap.add_argument('--cf-pairs',type=int,default=48); ap.add_argument('--device',default='auto'); ap.add_argument('--models',nargs='+',default=DEFAULT_MODELS); a=ap.parse_args() | |
| root=Path(__file__).resolve().parents[1]; out=root/'results/v10_cf_integrity'; out.mkdir(parents=True,exist_ok=True); data=root/'data/v9_train_family.npz' | |
| if not data.exists(): generate_family_npz(data) | |
| audit=build_oracle_audit(n=a.cf_pairs,horizon=a.cf_horizon,seed=10101); assert_audit_integrity(audit) | |
| with (out/'oracle_physics_audit.csv').open('w',newline='') as f: w=csv.DictWriter(f,fieldnames=list(audit[0])); w.writeheader(); w.writerows(audit) | |
| device=('cuda' if torch.cuda.is_available() else 'cpu') if a.device=='auto' else a.device; dev=torch.device(device); raw=[] | |
| for seed in a.seeds: | |
| for m in a.models: | |
| print(f'V10 seed={seed} model={m}'); sd=out/f'seed_{seed}' | |
| r=train_one_model(m,data,sd,epochs=a.epochs,batch_size=a.batch_size,seed=seed,device=device,cf_train_horizon=10,cf_pairs=48,cf_every=4,lambda_cf_dir=.02,lambda_cf_mag=.005,lambda_cf_branch=.08) | |
| ck=torch.load(sd/f'{m}.pt',map_location=dev); model=build_model(m).to(dev); model.load_state_dict(ck['model']); row={'seed':seed,'model':m,'params':r['params']} | |
| for rn,cfg in REGIMES.items(): | |
| met=evaluate_counterfactual_audit(model,dev,cfg,n=a.cf_pairs,horizon=a.cf_horizon,seed=10101) | |
| row.update({f'{rn}_{k}':v for k,v in met.items()}) | |
| raw.append(row) | |
| fields=[] | |
| for r in raw: | |
| for k in r: | |
| if k not in fields: fields.append(k) | |
| with (out/'metrics_by_seed.csv').open('w',newline='') as f: w=csv.DictWriter(f,fieldnames=fields); w.writeheader(); w.writerows(raw) | |
| summary=[] | |
| for m in a.models: | |
| rs=[r for r in raw if r['model']==m]; row={'model':m,'params':rs[0]['params'],'n_seeds':len(rs)} | |
| for k in fields: | |
| if k in ('seed','model','params') or not all(k in r and isinstance(r[k],(int,float)) for r in rs): continue | |
| mu,sd=stats([float(r[k]) for r in rs]); row[k+'_mean']=mu; row[k+'_std']=sd | |
| summary.append(row) | |
| sf=[] | |
| for r in summary: | |
| for k in r: | |
| if k not in sf: sf.append(k) | |
| with (out/'metrics_summary.csv').open('w',newline='') as f: w=csv.DictWriter(f,fieldnames=sf); w.writeheader(); w.writerows(summary) | |
| with (out/'run_config.json').open('w') as f: json.dump(vars(a)|{'device_resolved':device,'audit_seed':10101},f,indent=2) | |
| print('Oracle audit PASS; wrote',out/'metrics_summary.csv') | |
| if __name__=='__main__': main() | |