#!/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()