Download scripts/run_stability_benchmark.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 4.92 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_stability_benchmark.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_stability_benchmark.py
-
curl -L -o run_stability_benchmark.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_stability_benchmark.py
4.92 kB
| """Targeted attribution benchmark for SWM transition stability.""" | |
| from __future__ import annotations | |
| import argparse, csv, json | |
| from pathlib import Path | |
| from statistics import mean, stdev | |
| import torch | |
| from spectral_world_models.data import DatasetConfig, generate_benchmark_npz | |
| from spectral_world_models.train import train_one_model | |
| DEFAULT_MODELS=["koopman","neural_operator","swm","swm_selective","swm_spectral_norm","no_spectral_transition","no_stability_penalty"] | |
| BASE=[("test_psnr",lambda r:r["test"]["psnr"]),("test_ssim",lambda r:r["test"]["ssim"]),("test_nll",lambda r:r["test"]["nll"]),("test_mse",lambda r:r["test"]["mse"]),("test_text_acc",lambda r:r["test"]["text_acc"])] | |
| def stats(xs): return mean(xs), stdev(xs) if len(xs)>1 else 0.0 | |
| def main(): | |
| ap=argparse.ArgumentParser(description="SWM v4 benchmark: selective spectral regulation and long-horizon multimodal information preservation.") | |
| ap.add_argument("--data",default=None); ap.add_argument("--rollout-data",default=None); ap.add_argument("--out",default=None) | |
| 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("--horizons",type=int,nargs="+",default=[5,10,20,30,50,100]) | |
| ap.add_argument("--device",default="auto"); ap.add_argument("--models",nargs="+",default=DEFAULT_MODELS); ap.add_argument("--lambda-stability",type=float,default=0.01) | |
| a=ap.parse_args(); root=Path(__file__).resolve().parents[1] | |
| data=Path(a.data) if a.data else root/"data"/"mini_moving_shapes.npz" | |
| out=Path(a.out) if a.out else root/"results"/"stability_attribution" | |
| rollout=Path(a.rollout_data) if a.rollout_data else root/f"data/mini_moving_shapes_rollout_h{max(a.horizons)+1}.npz" | |
| if not data.exists(): generate_benchmark_npz(data,DatasetConfig()) | |
| need=max(a.horizons)+1; regen=True | |
| if rollout.exists(): | |
| try: | |
| import numpy as np | |
| with np.load(rollout,allow_pickle=True) as d: regen=d["test_images"].shape[1] < need | |
| except Exception: pass | |
| if regen: generate_benchmark_npz(rollout,DatasetConfig(seq_len=need,train_sequences=1,val_sequences=1,test_sequences=64,seed=7007)) | |
| device=("cuda" if torch.cuda.is_available() else "cpu") if a.device=="auto" else a.device | |
| if device.startswith("cuda") and not torch.cuda.is_available(): raise RuntimeError("CUDA requested but unavailable") | |
| out.mkdir(parents=True,exist_ok=True); grouped={m:[] for m in a.models}; raw=[] | |
| for seed in a.seeds: | |
| for m in a.models: | |
| print(f"\n=== seed={seed} model={m} horizons={a.horizons} ===") | |
| r=train_one_model(m,data,out/f"seed_{seed}",epochs=a.epochs,batch_size=a.batch_size,seed=seed,device=device,rollout_data_path=rollout,rollout_horizons=tuple(a.horizons),lambda_stability=a.lambda_stability) | |
| grouped[m].append(r); row={"seed":seed,"model":m,"params":r["params"]} | |
| for k,g in BASE: row[k]=g(r) | |
| for h in a.horizons: | |
| rr=r["rollout_by_horizon"][str(h)] | |
| for metric in ["rollout_mse","rollout_text_acc","rollout_latent_cosine","rollout_joint_score"]: row[f"{metric}_h{h}"]=rr[metric] | |
| for k,v in r.get("stability_diagnostics",{}).items(): row[k]=v | |
| 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,runs in grouped.items(): | |
| row={"model":m,"params":runs[0]["params"],"n_seeds":len(runs)} | |
| for k,g in BASE: | |
| mu,sd=stats([float(g(r)) for r in runs]); row[k+"_mean"]=mu; row[k+"_std"]=sd | |
| for h in a.horizons: | |
| for metric in ["rollout_mse","rollout_text_acc","rollout_latent_cosine","rollout_joint_score"]: | |
| vals=[float(r["rollout_by_horizon"][str(h)][metric]) for r in runs]; mu,sd=stats(vals); row[f"{metric}_h{h}_mean"]=mu; row[f"{metric}_h{h}_std"]=sd | |
| diag_keys=set().union(*(r.get("stability_diagnostics",{}).keys() for r in runs)) | |
| for k in sorted(diag_keys): | |
| vals=[float(r["stability_diagnostics"][k]) for r in runs if k in r.get("stability_diagnostics",{})]; mu,sd=stats(vals); row[k+"_mean"]=mu; row[k+"_std"]=sd | |
| summary.append(row) | |
| fields=[] | |
| for r in summary: | |
| for k in r: | |
| if k not in fields: fields.append(k) | |
| with (out/"metrics_summary.csv").open("w",newline="") as f: | |
| w=csv.DictWriter(f,fieldnames=fields); w.writeheader(); w.writerows(summary) | |
| with (out/"run_config.json").open("w") as f: json.dump(vars(a)|{"device_resolved":device,"rollout_data_resolved":str(rollout)},f,indent=2) | |
| print("\nWrote",out/"metrics_summary.csv") | |
| if __name__=="__main__": main() | |