kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
Raw History Blame Contribute Delete
5.47 kB
"""V8 frozen-architecture OOD generalization + causal-control + planning 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.hard_dynamics import HardDynamicsConfig,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.counterfactual import evaluate_counterfactual
from spectral_world_models.models import build_model
from spectral_world_models.v8_generalization import OOD_CONFIG,evaluate_counterfactual_cfg,evaluate_planning
DEFAULT_MODELS=['neural_operator','swm_selective','swm_structured_cf','no_spectral_transition']
def stats(x): return mean(x),stdev(x) if len(x)>1 else 0.
def main():
ap=argparse.ArgumentParser(description='V8: freeze V7 architectures; test ID/OOD generalization, intervention, and planning.')
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=96); ap.add_argument('--device',default='auto'); ap.add_argument('--models',nargs='+',default=DEFAULT_MODELS)
ap.add_argument('--cf-train-horizon',type=int,default=8); ap.add_argument('--cf-pairs',type=int,default=64); 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=.10); a=ap.parse_args()
root=Path(__file__).resolve().parents[1]; train_data=root/'data/hard_controlled_dynamics_v6.npz'; id_roll=root/f'data/v8_id_rollout_h{max(a.horizons)+1}.npz'; ood=root/f'data/v8_ood_rollout_h{max(a.horizons)+1}.npz'; out=root/'results/v8_generalization'; out.mkdir(parents=True,exist_ok=True)
if not train_data.exists(): generate_hard_benchmark_npz(train_data,HardDynamicsConfig())
if not id_roll.exists(): generate_hard_benchmark_npz(id_roll,HardDynamicsConfig(seq_len=max(a.horizons)+1,train_sequences=1,val_sequences=1,test_sequences=96,seed=8016))
oc=OOD_CONFIG; oc.seq_len=max(a.horizons)+1
if not ood.exists(): generate_hard_benchmark_npz(ood,oc)
device=('cuda' if torch.cuda.is_available() else 'cpu') if a.device=='auto' else a.device; raw=[]
for seed in a.seeds:
for m in a.models:
print(f'\n=== V8 seed={seed} model={m} ==='); sd=out/f'seed_{seed}'
r=train_one_model(m,train_data,sd,epochs=a.epochs,batch_size=a.batch_size,seed=seed,device=device,rollout_data_path=id_roll,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=device); model=build_model(m).to(device); model.load_state_dict(ck['model'])
ood_t=evaluate_transitions(model,DataLoader(TransitionDataset(ood,'test'),batch_size=a.batch_size),torch.device(device)); ood_s=DataLoader(SequenceDataset(ood,'test'),batch_size=a.batch_size)
row={'seed':seed,'model':m,'params':r['params'],'id_test_mse':r['test']['mse'],'id_test_text_acc':r['test']['text_acc'],'ood_test_mse':ood_t['mse'],'ood_test_text_acc':ood_t['text_acc']}
idcf=evaluate_counterfactual(model,torch.device(device),horizon=a.cf_horizon,seed=8606+seed); row.update({'id_'+k:v for k,v in idcf.items()})
row.update(evaluate_counterfactual_cfg(model,torch.device(device),horizon=a.cf_horizon,seed=8809+seed,cfg=oc))
for h in a.horizons:
rr=r['rollout_by_horizon'][str(h)]; oo=evaluate_rollout(model,ood_s,torch.device(device),horizon=h)
for q in ['rollout_mse','rollout_text_acc','rollout_latent_cosine']: row[f'id_{q}_h{h}']=rr[q]; row[f'ood_{q}_h{h}']=oo[q]
for prefix,cfg in [('id',HardDynamicsConfig()),('ood',oc)]:
p=evaluate_planning(model,torch.device(device),n=a.planning_episodes,horizon=a.planning_horizon,seed=8900+seed,cfg=cfg); row.update({prefix+'_'+k:v for k,v in p.items()})
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,'architecture_frozen_from':'v7','ood_config':oc.__dict__},f,indent=2)
print('\nWrote',out/'metrics_summary.csv')
if __name__=='__main__':main()