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