Download scripts/run_v13_axis_audit.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 3.88 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v13_axis_audit.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_v13_axis_audit.py
-
curl -L -o run_v13_axis_audit.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v13_axis_audit.py
3.88 kB
| from __future__ import annotations | |
| import argparse,csv,json | |
| from pathlib import Path | |
| import numpy as np, 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.v13_axis_equivariance import REGIMES,audit_axis_equivariance,audit_forced_collisions | |
| 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); w.writeheader(); w.writerows(rows) | |
| def summarize(rows,keys,metrics): | |
| out=[]; groups={} | |
| for r in rows: groups.setdefault(tuple(r[k] for k in keys),[]).append(r) | |
| for kval,rs in groups.items(): | |
| z=dict(zip(keys,kval)); z['n']=len(rs) | |
| for m in metrics: | |
| x=np.asarray([float(q[m]) for q in rs if np.isfinite(float(q[m]))]); z[m+'_mean']=float(x.mean()) if len(x) else float('nan'); z[m+'_std']=float(x.std(ddof=1)) if len(x)>1 else 0. | |
| out.append(z) | |
| return out | |
| def main(): | |
| ap=argparse.ArgumentParser(description='V13 frozen-V9 axis symmetry, intervention equivariance, and forced-collision audit') | |
| 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('--pairs',type=int,default=64); ap.add_argument('--horizon',type=int,default=30); ap.add_argument('--collision-horizon',type=int,default=15); 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/v13_axis_equivariance'; 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); axis=[]; coll=[] | |
| for seed in a.seeds: | |
| for m in a.models: | |
| print(f'V13 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']) | |
| for rn,cfg in REGIMES.items(): | |
| rr=audit_axis_equivariance(model,dev,cfg,rn,n=a.pairs,horizon=a.horizon,seed=13101); cc=audit_forced_collisions(model,dev,cfg,rn,n=a.pairs,horizon=a.collision_horizon,seed=13201) | |
| for x in rr:x.update(seed=seed,model=m,params=r['params']) | |
| for x in cc:x.update(seed=seed,model=m,params=r['params']) | |
| axis.extend(rr); coll.extend(cc) | |
| writecsv(out/'axis_equivariance_by_episode.csv',axis); writecsv(out/'forced_collision_by_episode.csv',coll) | |
| axsum=summarize(axis,['model','regime','program','transform'],['base_cosine','transformed_cosine','cosine_drop','base_magnitude_ratio','transformed_magnitude_ratio','equivariance_effect_cosine','equivariance_relative_error','oracle_equivariance_error']) | |
| csum=summarize(coll,['model','regime','wall','program'],['trajectory_cosine','magnitude_ratio','boundary_contact']) | |
| writecsv(out/'axis_equivariance_summary.csv',axsum); writecsv(out/'forced_collision_summary.csv',csum) | |
| with (out/'run_config.json').open('w') as f: json.dump(vars(a)|{'release':'V13 Axis Symmetry & Intervention Equivariance Audit','learning_changes':'none; V9 recipe frozen','device_resolved':device},f,indent=2) | |
| print('wrote',out) | |
| if __name__=='__main__':main() | |