Download scripts/run_v14_external_benchmark.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 5.04 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v14_external_benchmark.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_v14_external_benchmark.py
-
curl -L -o run_v14_external_benchmark.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v14_external_benchmark.py
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() | |