kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
Raw History Blame Contribute Delete
2.96 kB
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.v11_causal_calibration import REGIMES,evaluate_causal_calibration
DEFAULT_MODELS=['neural_operator','swm_structured_cf_v9','swm_structured_cf_v11','no_spectral_transition']
def stats(x): return mean(x),stdev(x) if len(x)>1 else 0.
def main():
ap=argparse.ArgumentParser(description='V11 causal-gain calibration benchmark')
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('--device',default='auto'); ap.add_argument('--models',nargs='+',default=DEFAULT_MODELS); a=ap.parse_args()
root=Path(__file__).resolve().parents[1]; out=root/'results/v11_causal_calibration'; 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); raw=[]
for seed in a.seeds:
for m in a.models:
print(f'V11 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=12,cf_pairs=a.cf_pairs,cf_every=4,lambda_cf_dir=.03,lambda_cf_mag=.01,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(): row.update({f'{rn}_{k}':v for k,v in evaluate_causal_calibration(model,dev,cfg,n=a.cf_pairs,horizon=a.cf_horizon,seed=11101).items()})
raw.append(row)
fields=list(raw[0]);
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'): continue
mu,sd=stats([float(r[k]) for r in rs]); row[k+'_mean']=mu; row[k+'_std']=sd
summary.append(row)
sf=list(summary[0]);
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,'architecture_base':'v9 structured transport','objective':'scale-aware trajectory causal calibration'},f,indent=2)
print('wrote',out/'metrics_summary.csv')
if __name__=='__main__': main()