Download scripts/run_v15_external_diagnostics.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 5.74 kB
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v15_external_diagnostics.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/scripts/run_v15_external_diagnostics.py
-
curl -L -o run_v15_external_diagnostics.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/scripts/run_v15_external_diagnostics.py
5.74 kB
| #!/usr/bin/env python3 | |
| """V15.1: bug-fixed diagnostic-only evaluation of frozen V14 checkpoints. No retraining. | |
| Fixes V15 token/state diagnostic shape handling by using WorldModel.decode_text(), | |
| which reshapes the flat decoder output to [batch, token_position, vocabulary]. | |
| """ | |
| from __future__ import annotations | |
| import argparse,csv,json,sys | |
| from pathlib import Path | |
| import numpy as np, torch | |
| import torch.nn.functional as F | |
| 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 ExternalSequenceDataset,load_meta,SUPPORTED_ENVS | |
| from torch.utils.data import DataLoader | |
| MODELS=("neural_operator","swm_structured_cf_v9","no_spectral_transition") | |
| def ssim_simple(x,y): | |
| # global SSIM per image, sufficient as a dependency-free diagnostic | |
| mu_x=x.flatten(1).mean(1); mu_y=y.flatten(1).mean(1) | |
| vx=((x.flatten(1)-mu_x[:,None])**2).mean(1); vy=((y.flatten(1)-mu_y[:,None])**2).mean(1) | |
| cov=((x.flatten(1)-mu_x[:,None])*(y.flatten(1)-mu_y[:,None])).mean(1) | |
| return ((2*mu_x*mu_y+1e-4)*(2*cov+9e-4)/((mu_x**2+mu_y**2+1e-4)*(vx+vy+9e-4))).mean().item() | |
| def diag(model,loader,device,horizons): | |
| acc={h:{k:[] for k in ['mse','static_mse','fg_mse','motion_mse','psnr','ssim','latent_cos','token_acc']} 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) | |
| z0=model.encode(imgs[:,0],toks[:,0]); z=z0; prev_pred=imgs[:,0] | |
| preds=[]; zs=[] | |
| for t in range(acts.shape[1]): | |
| z=model.transition(z,acts[:,t]); preds.append(model.image_decoder(z)); zs.append(z) | |
| for h in horizons: | |
| t=min(h,len(preds))-1; pred=preds[t]; gt=imgs[:,t+1]; base=imgs[:,0] | |
| err=(pred-gt)**2; mse=err.flatten(1).mean(1) | |
| static=((base-gt)**2).flatten(1).mean(1) | |
| mask=(gt-base).abs()>0.03; fg=(err*mask).flatten(1).sum(1)/mask.flatten(1).sum(1).clamp_min(1) | |
| if t>0: | |
| pm=pred-preds[t-1]; gm=gt-imgs[:,t] | |
| else: pm=pred-base; gm=gt-base | |
| motion=((pm-gm)**2).flatten(1).mean(1) | |
| truez=model.encode(gt,toks[:,t+1]); cos=F.cosine_similarity(zs[t],truez,dim=1) | |
| logits=model.decode_text(zs[t]) | |
| target=toks[:,t+1] | |
| if logits.ndim != 3 or logits.shape[:2] != target.shape: | |
| raise RuntimeError(f'Incompatible state-token shapes: logits={tuple(logits.shape)}, target={tuple(target.shape)}. Expected [B,L,V] versus [B,L].') | |
| predtok=logits.argmax(-1); valid=target!=0 | |
| tok=((predtok==target)&valid).sum(1)/valid.sum(1).clamp_min(1) | |
| acc[h]['mse'] += mse.cpu().tolist(); acc[h]['static_mse'] += static.cpu().tolist(); acc[h]['fg_mse'] += fg.cpu().tolist(); acc[h]['motion_mse'] += motion.cpu().tolist(); acc[h]['psnr'] += (-10*torch.log10(mse.clamp_min(1e-10))).cpu().tolist(); acc[h]['ssim'].append(ssim_simple(pred,gt)); acc[h]['latent_cos'] += cos.cpu().tolist(); acc[h]['token_acc'] += tok.cpu().tolist() | |
| out={} | |
| for h,d in acc.items(): | |
| for k,v in d.items(): out[f'{k}_h{h}']=float(np.mean(v)) | |
| return out | |
| def spectral_diag(model): | |
| d={} | |
| try: | |
| sd=model.stability_diagnostics() | |
| for k,v in sd.items(): d['spectral_'+k]=float(v.detach().cpu() if torch.is_tensor(v) else v) | |
| except Exception: pass | |
| return d | |
| def main(): | |
| p=argparse.ArgumentParser(description='V15.1 bug-fixed external dynamics diagnostic; evaluates frozen V14 checkpoints without retraining.'); p.add_argument('--v14-results',default='results/v14_external_control'); 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=list(range(10))); p.add_argument('--horizons',nargs='+',type=int,default=[5,10,20,30]); p.add_argument('--batch-size',type=int,default=64); p.add_argument('--device',default='cpu'); p.add_argument('--out',default='results/v15_external_diagnostics'); a=p.parse_args() | |
| root=Path(a.v14_results); out=Path(a.out); out.mkdir(parents=True,exist_ok=True); rows=[] | |
| for e in a.envs: | |
| data=root/'data'/f'{e}.npz'; meta=load_meta(data); loader=DataLoader(ExternalSequenceDataset(data,'test'),batch_size=a.batch_size) | |
| for m in a.models: | |
| for s in a.seeds: | |
| ck=root/'checkpoints'/e/f'{m}_seed{s}.pt'; print('[V15]',e,m,s,flush=True) | |
| model=build_model(m,image_size=32,vocab_size=meta['vocab_size'],max_text_len=8,action_dim=meta['action_dim']).to(a.device); model.load_state_dict(torch.load(ck,map_location=a.device)['model']) | |
| r={'env_id':e,'model':m,'seed':s}; r.update(diag(model,loader,torch.device(a.device),a.horizons)); r.update(spectral_diag(model)); rows.append(r) | |
| fields=sorted(set().union(*(r.keys() for r in rows))); | |
| with open(out/'diagnostics_by_seed.csv','w',newline='') as f: w=csv.DictWriter(f,fieldnames=fields); w.writeheader(); w.writerows(rows) | |
| 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]; q={'env_id':e,'model':m,'n_seeds':len(rr)} | |
| for k in fields: | |
| if k in ('env_id','model','seed') or not all(k in x for x in rr): continue | |
| vals=np.array([x[k] for x in rr],float); q[k+'_mean']=vals.mean(); q[k+'_std']=vals.std(ddof=1) if len(vals)>1 else 0 | |
| summary.append(q) | |
| sf=sorted(set().union(*(r.keys() for r in summary))); | |
| with open(out/'diagnostics_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/'diagnostics_summary.csv') | |
| if __name__=='__main__': main() | |