Download scripts/run_v12_decomposition.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 5.29 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v12_decomposition.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_v12_decomposition.py
-
curl -L -o run_v12_decomposition.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v12_decomposition.py
5.29 kB
| from __future__ import annotations | |
| import argparse,csv,json | |
| from pathlib import Path | |
| 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.v12_causal_decomposition import REGIMES,decompose_model,summarize_strata,paired_bootstrap_delta | |
| DEFAULT_MODELS=['neural_operator','swm_structured_cf_v9','no_spectral_transition'] | |
| def writecsv(path,rows): | |
| if not rows:return | |
| fields=[] | |
| for r in rows: | |
| for k in r: | |
| if k not in fields: fields.append(k) | |
| with path.open('w',newline='') as f: | |
| w=csv.DictWriter(f,fieldnames=fields,extrasaction='ignore'); w.writeheader(); w.writerows(rows) | |
| def main(): | |
| ap=argparse.ArgumentParser(description='V12 causal-response decomposition; frozen V9 architecture/objective') | |
| 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=list(range(10))) | |
| ap.add_argument('--cf-horizon',type=int,default=30); ap.add_argument('--cf-pairs',type=int,default=64); ap.add_argument('--bootstrap',type=int,default=2000); 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/v12_causal_decomposition'; out.mkdir(parents=True,exist_ok=True); data=root/'data/v9_train_family.npz' | |
| if not data.exists(): generate_family_npz(data) | |
| device=('cuda' if torch.cuda.is_available() else 'cpu') if a.device=='auto' else a.device; dev=torch.device(device); allrows=[] | |
| for seed in a.seeds: | |
| for m in a.models: | |
| print(f'V12 seed={seed} model={m}'); sd=out/f'seed_{seed}' | |
| # Freeze the V9 experimental recipe: V12 changes evaluation, not learning. | |
| 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']) | |
| for rn,cfg in REGIMES.items(): | |
| rr=decompose_model(model,dev,cfg,rn,n=a.cf_pairs,horizon=a.cf_horizon,seed=12101) | |
| for x in rr:x.update({'seed':seed,'model':m,'params':r['params']}) | |
| allrows.extend(rr) | |
| writecsv(out/'decomposition_by_episode.csv',allrows) | |
| summaries=[] | |
| for seed in a.seeds: | |
| for m in a.models: | |
| for rn in REGIMES: | |
| rs=[r for r in allrows if r['seed']==seed and r['model']==m and r['regime']==rn] | |
| for s in summarize_strata(rs): s.update({'seed':seed,'model':m,'regime':rn}); summaries.append(s) | |
| writecsv(out/'stratified_metrics_by_seed.csv',summaries) | |
| # Aggregate across seeds as observations; preserve seed-level file for uncertainty inspection. | |
| agg=[] | |
| keys=sorted({(r['model'],r['regime'],r['stratum_type'],str(r.get('program','')),str(r.get('horizon','')),str(r.get('effect_bin','')),str(r.get('speed_bin','')),str(r.get('boundary_contact','')),str(r.get('occlusion_zone_exposure',''))) for r in summaries}) | |
| for k in keys: | |
| rs=[r for r in summaries if (r['model'],r['regime'],r['stratum_type'],str(r.get('program','')),str(r.get('horizon','')),str(r.get('effect_bin','')),str(r.get('speed_bin','')),str(r.get('boundary_contact','')),str(r.get('occlusion_zone_exposure','')))==k] | |
| row={'model':k[0],'regime':k[1],'stratum_type':k[2],'program':k[3],'horizon':k[4],'effect_bin':k[5],'speed_bin':k[6],'boundary_contact':k[7],'occlusion_zone_exposure':k[8],'n_seeds':len(rs)} | |
| for metric in ('trajectory_cosine','magnitude_ratio','magnitude_ratio_mae','log_gain_error'): | |
| vals=[float(x[metric+'_mean']) for x in rs]; row[metric+'_mean']=sum(vals)/len(vals); row[metric+'_seed_std']=float(torch.tensor(vals).std(unbiased=True)) if len(vals)>1 else 0. | |
| agg.append(row) | |
| writecsv(out/'stratified_metrics_summary.csv',agg) | |
| # Paired, episode-clustered CIs within each seed/regime, then pooled matched rows across seeds via seed-offset episode IDs. | |
| boots=[]; v9='swm_structured_cf_v9' | |
| for control in [m for m in a.models if m!=v9]: | |
| for rn in REGIMES: | |
| A=[];B=[] | |
| for si,seed in enumerate(a.seeds): | |
| ar=[dict(r,episode=r['episode']+si*100000) for r in allrows if r['model']==v9 and r['regime']==rn and r['seed']==seed] | |
| br=[dict(r,episode=r['episode']+si*100000) for r in allrows if r['model']==control and r['regime']==rn and r['seed']==seed] | |
| A.extend(ar);B.extend(br) | |
| for metric in ('trajectory_cosine','magnitude_ratio_mae','log_gain_error'): | |
| z=paired_bootstrap_delta(A,B,metric=metric,n_boot=a.bootstrap,seed=12201); z.update({'regime':rn,'control':control}); boots.append(z) | |
| writecsv(out/'paired_bootstrap_v9_vs_controls.csv',boots) | |
| with (out/'run_config.json').open('w') as f: json.dump(vars(a)|{'device_resolved':device,'release':'V12 Causal Response Decomposition','learning_changes':'none; V9 recipe frozen','primary_model':'swm_structured_cf_v9'},f,indent=2) | |
| print('wrote',out) | |
| if __name__=='__main__': main() | |