"""V9: multi-environment trajectory-CF causal generalization benchmark.""" from __future__ import annotations import argparse,csv,json from pathlib import Path from statistics import mean,stdev import torch from torch.utils.data import DataLoader from spectral_world_models.v9_causal_generalization import generate_family_npz,INTERP,EXTRAP,STRUCTURAL from spectral_world_models.hard_dynamics import generate_hard_benchmark_npz from spectral_world_models.data import TransitionDataset,SequenceDataset from spectral_world_models.train import train_one_model,evaluate_transitions,evaluate_rollout from spectral_world_models.v8_generalization import evaluate_counterfactual_cfg,evaluate_planning from spectral_world_models.models import build_model 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='V9 causal generalization: environment-family training + trajectory counterfactual loss.') 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('--cf-horizon',type=int,default=30); ap.add_argument('--planning-horizon',type=int,default=20); ap.add_argument('--planning-episodes',type=int,default=64); ap.add_argument('--device',default='auto'); ap.add_argument('--models',nargs='+',default=DEFAULT_MODELS); ap.add_argument('--cf-train-horizon',type=int,default=10); ap.add_argument('--cf-pairs',type=int,default=48); ap.add_argument('--cf-every',type=int,default=4); ap.add_argument('--lambda-cf-dir',type=float,default=.02); ap.add_argument('--lambda-cf-mag',type=float,default=.005); ap.add_argument('--lambda-cf-branch',type=float,default=.08); a=ap.parse_args() root=Path(__file__).resolve().parents[1]; data=root/'data/v9_train_family.npz'; out=root/'results/v9_causal_generalization'; out.mkdir(parents=True,exist_ok=True); maxh=max(a.horizons) if not data.exists(): generate_family_npz(data) regimes={'interpolation':INTERP,'extrapolation':EXTRAP,'structural_ood':STRUCTURAL}; paths={} for name,cfg in regimes.items(): cfg.seq_len=maxh+1; p=root/f'data/v9_{name}_h{maxh+1}.npz'; paths[name]=p if not p.exists(): generate_hard_benchmark_npz(p,cfg) 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'\n=== V9 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,rollout_data_path=paths['interpolation'],rollout_horizons=tuple(a.horizons),cf_train_horizon=a.cf_train_horizon,cf_pairs=a.cf_pairs,cf_every=a.cf_every,lambda_cf_dir=a.lambda_cf_dir,lambda_cf_mag=a.lambda_cf_mag,lambda_cf_branch=a.lambda_cf_branch) 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 name,cfg in regimes.items(): p=paths[name]; tm=evaluate_transitions(model,DataLoader(TransitionDataset(p,'test'),batch_size=a.batch_size),dev); row[f'{name}_test_mse']=tm['mse']; row[f'{name}_test_text_acc']=tm['text_acc'] sl=DataLoader(SequenceDataset(p,'test'),batch_size=a.batch_size) for h in a.horizons: rr=evaluate_rollout(model,sl,dev,horizon=h) for q in ('rollout_mse','rollout_text_acc','rollout_latent_cosine'): row[f'{name}_{q}_h{h}']=rr[q] cf=evaluate_counterfactual_cfg(model,dev,horizon=a.cf_horizon,n=96,seed=9300+seed,cfg=cfg) for k,v in cf.items(): row[f'{name}_{k.removeprefix("ood_")}']=v pl=evaluate_planning(model,dev,n=a.planning_episodes,horizon=a.planning_horizon,seed=9400+seed,cfg=cfg) for k,v in pl.items(): row[f'{name}_{k}']=v row.update(r.get('stability_diagnostics',{})); 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,'regimes':{k:v.__dict__ for k,v in regimes.items()}},f,indent=2) print('\nWrote',out/'metrics_summary.csv') if __name__=='__main__': main()