kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
Raw History Blame Contribute Delete
4.32 kB
"""V6 structured spectral transport + counterfactual action benchmark."""
from __future__ import annotations
import argparse,csv,json
from pathlib import Path
from statistics import mean,stdev
import torch, numpy as np
from spectral_world_models.hard_dynamics import HardDynamicsConfig,generate_hard_benchmark_npz
from spectral_world_models.train import train_one_model
from spectral_world_models.counterfactual import evaluate_counterfactual
from spectral_world_models.models import build_model
DEFAULT_MODELS=['neural_operator','swm','swm_selective','swm_structured','no_spectral_transition','no_stability_penalty']
def stats(x): return mean(x),stdev(x) if len(x)>1 else 0.
def main():
ap=argparse.ArgumentParser(); 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('--device',default='auto'); ap.add_argument('--models',nargs='+',default=DEFAULT_MODELS); a=ap.parse_args()
root=Path(__file__).resolve().parents[1]; data=root/'data/hard_controlled_dynamics_v6.npz'; rollout=root/f'data/hard_controlled_dynamics_v6_rollout_h{max(a.horizons)+1}.npz'; out=root/'results/v6_structured_transport'; out.mkdir(parents=True,exist_ok=True)
if not data.exists(): generate_hard_benchmark_npz(data,HardDynamicsConfig())
if not rollout.exists(): generate_hard_benchmark_npz(rollout,HardDynamicsConfig(seq_len=max(a.horizons)+1,train_sequences=1,val_sequences=1,test_sequences=96,seed=7017))
device=('cuda' if torch.cuda.is_available() else 'cpu') if a.device=='auto' else a.device; 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} ==='); 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=rollout,rollout_horizons=tuple(a.horizons))
ck=torch.load(sd/f'{m}.pt',map_location=device); model=build_model(m).to(device); model.load_state_dict(ck['model']); cf=evaluate_counterfactual(model,torch.device(device),horizon=a.cf_horizon,seed=606+seed); r['counterfactual']=cf; grouped[m].append(r)
row={'seed':seed,'model':m,'params':r['params'],'test_mse':r['test']['mse'],'test_text_acc':r['test']['text_acc'],**cf}
for h in a.horizons:
rr=r['rollout_by_horizon'][str(h)]; row[f'rollout_mse_h{h}']=rr['rollout_mse']; row[f'rollout_text_acc_h{h}']=rr['rollout_text_acc']; row[f'rollout_latent_cosine_h{h}']=rr['rollout_latent_cosine']
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,runs in grouped.items():
row={'model':m,'params':runs[0]['params'],'n_seeds':len(runs)}
getters={'test_mse':lambda r:r['test']['mse'],'test_text_acc':lambda r:r['test']['text_acc'],'cf_final_image_mse':lambda r:r['counterfactual']['cf_final_image_mse'],'cf_effect_vector_cosine':lambda r:r['counterfactual']['cf_effect_vector_cosine'],'cf_effect_magnitude_ratio':lambda r:r['counterfactual']['cf_effect_magnitude_ratio']}
for h in a.horizons:
for q in ['rollout_mse','rollout_text_acc','rollout_latent_cosine']: getters[f'{q}_h{h}']=lambda r,q=q,h=h:r['rollout_by_horizon'][str(h)][q]
for k,g in getters.items(): mu,sd=stats([float(g(r)) for r in runs]);row[k+'_mean']=mu;row[k+'_std']=sd
dkeys=set().union(*(r.get('stability_diagnostics',{}).keys() for r in runs))
for k in dkeys:
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},f,indent=2)
print('\nWrote',out/'metrics_summary.csv')
if __name__=='__main__':main()