Spectral-World-Models-Reproducibility / scripts /run_v14_external_benchmark.py
kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
Raw History Blame Contribute Delete
5.04 kB
#!/usr/bin/env python3
"""Train frozen architecture families from scratch on external Gymnasium controls."""
from __future__ import annotations
import argparse,csv,json,sys
from pathlib import Path
import numpy as np, torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
sys.path.insert(0,str(Path(__file__).resolve().parents[1]))
from spectral_world_models.models import build_model
from spectral_world_models.v14_external_control import ExternalConfig,collect_external_npz,ExternalTransitionDataset,ExternalSequenceDataset,load_meta,SUPPORTED_ENVS
from spectral_world_models.metrics import count_parameters
MODELS=("neural_operator","swm_structured_cf_v9","no_spectral_transition")
def loss_fn(model,b,vocab_size):
o=model(b['image'],b['tokens'],b['action'])
with torch.no_grad(): tz=model.encode(b['next_image'],b['next_tokens'])
im=F.mse_loss(o['image_pred'],b['next_image']); lat=F.mse_loss(o['z_pred'],tz)
txt=F.cross_entropy(o['text_logits'].reshape(-1,vocab_size),b['next_tokens'].reshape(-1),ignore_index=0)
return im+.15*lat+.05*txt+.01*model.stability_penalty()
def rollout_metrics(model,loader,device,horizons):
sums={h:[] for h in horizons}
model.eval()
with torch.no_grad():
for b in loader:
imgs=b['images'].to(device); toks=b['tokens'].to(device); acts=b['actions'].to(device); z=model.encode(imgs[:,0],toks[:,0])
step=[]
for t in range(acts.shape[1]):
z=model.transition(z,acts[:,t]); p=model.image_decoder(z); step.append(((p-imgs[:,t+1])**2).flatten(1).mean(1).cpu())
for h in horizons:
hh=min(h,len(step)); sums[h].append(torch.stack(step[:hh],1).mean(1))
return {f'rollout_mse_h{h}':float(torch.cat(sums[h]).mean()) for h in horizons}
def train(env_id,model_name,seed,args,out):
data=out/'data'/f'{env_id}.npz'
if not data.exists():
collect_external_npz(data,ExternalConfig(env_id=env_id,seq_len=max(args.horizons)+1,train_episodes=args.train_episodes,val_episodes=args.val_episodes,test_episodes=args.test_episodes,seed=args.data_seed))
meta=load_meta(data); device=torch.device(args.device)
torch.manual_seed(seed)
model=build_model(model_name,image_size=32,vocab_size=meta['vocab_size'],max_text_len=8,action_dim=meta['action_dim']).to(device)
tr=DataLoader(ExternalTransitionDataset(data,'train'),batch_size=args.batch_size,shuffle=True); te=DataLoader(ExternalSequenceDataset(data,'test'),batch_size=args.batch_size)
opt=torch.optim.AdamW(model.parameters(),lr=args.lr)
model.train()
for _ in range(args.epochs):
for b in tr:
b={k:v.to(device) for k,v in b.items()}; opt.zero_grad(set_to_none=True); L=loss_fn(model,b,meta['vocab_size']); L.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),1.0); opt.step()
m=rollout_metrics(model,te,device,args.horizons)
m.update(env_id=env_id,model=model_name,seed=seed,params=count_parameters(model))
ck=out/'checkpoints'/env_id; ck.mkdir(parents=True,exist_ok=True); torch.save({'model':model.state_dict(),'meta':meta,'model_name':model_name,'seed':seed},ck/f'{model_name}_seed{seed}.pt')
return m
def main():
p=argparse.ArgumentParser(); p.add_argument('--envs',nargs='+',default=list(SUPPORTED_ENVS)); p.add_argument('--models',nargs='+',default=list(MODELS)); p.add_argument('--seeds',nargs='+',type=int,default=[0,1,2,3,4]); p.add_argument('--horizons',nargs='+',type=int,default=[5,10,20,30]); p.add_argument('--epochs',type=int,default=5); p.add_argument('--batch-size',type=int,default=64); p.add_argument('--lr',type=float,default=2e-3); p.add_argument('--train-episodes',type=int,default=256); p.add_argument('--val-episodes',type=int,default=64); p.add_argument('--test-episodes',type=int,default=64); p.add_argument('--data-seed',type=int,default=1400); p.add_argument('--device',default='cpu'); p.add_argument('--out',default='results/v14_external_control'); a=p.parse_args()
out=Path(a.out); out.mkdir(parents=True,exist_ok=True); rows=[]
for e in a.envs:
for m in a.models:
for s in a.seeds:
print(f'[V14] {e} {m} seed={s}',flush=True); rows.append(train(e,m,s,a,out))
fields=list(rows[0]);
with open(out/'metrics_by_seed.csv','w',newline='') as f: w=csv.DictWriter(f,fieldnames=fields); w.writeheader(); w.writerows(rows)
# aggregate mean/std by env/model
summary=[]
for e in a.envs:
for m in a.models:
rr=[r for r in rows if r['env_id']==e and r['model']==m]; row={'env_id':e,'model':m,'params':rr[0]['params'],'n_seeds':len(rr)}
for h in a.horizons:
vals=np.array([r[f'rollout_mse_h{h}'] for r in rr]); row[f'rollout_mse_h{h}_mean']=vals.mean(); row[f'rollout_mse_h{h}_std']=vals.std(ddof=1) if len(vals)>1 else 0.
summary.append(row)
sf=list(summary[0]);
with open(out/'metrics_summary.csv','w',newline='') as f: w=csv.DictWriter(f,fieldnames=sf); w.writeheader(); w.writerows(summary)
(out/'run_config.json').write_text(json.dumps(vars(a),indent=2)); print(out/'metrics_summary.csv')
if __name__=='__main__': main()