Spectral-World-Models-Reproducibility / scripts /run_v15_external_diagnostics.py
kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
Raw History Blame Contribute Delete
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()