"""Memory-bounded, resumable evaluation using official LeWM task/CEM configs.""" import os os.environ.setdefault('MUJOCO_GL','egl') os.environ.setdefault('SDL_VIDEODRIVER','dummy') import sys, json, time, argparse, hashlib, csv from pathlib import Path ROOT=Path(__file__).resolve().parents[1] sys.path.insert(0,str(ROOT/'runtime')) import numpy as np, torch, h5py from omegaconf import OmegaConf from sklearn.preprocessing import StandardScaler from torchvision.transforms import v2 as T import stable_worldmodel as swm from reporting import save_sheet, result_path FILES={'tworoom':'tworoom.h5','cube':'cube_single_expert.h5','pusht':'pusht_expert_train.h5'} def manifest(task,offset,seed=42): p=ROOT/'manifests'/(f'{task}_{offset}.json' if seed==42 else f'{task}_{offset}_seed{seed}.json') if p.exists(): return json.loads(p.read_text()) with h5py.File('/workspace/datasets/'+FILES[task]) as f: col='ep_idx' if 'ep_idx' in f else 'episode_idx'; ep=f[col][:];steps=f['step_idx'][:] ids,starts=np.unique(ep,return_index=True); lens=np.maximum.reduceat(steps,starts)+1 valid=np.flatnonzero(steps<=np.repeat(lens,np.diff(np.r_[starts,len(ep)]))-offset-1) # Preserve the upstream sampling convention (last valid row is excluded). chosen=np.sort(valid[np.random.default_rng(seed).choice(len(valid)-1,200,replace=False)]) d={'task':task,'offset':offset,'seed':seed,'dataset':str(Path('/workspace/datasets')/FILES[task]),'rows':chosen.tolist(),'episodes':ep[chosen].tolist(),'start_steps':steps[chosen].tolist()} p.write_text(json.dumps(d,indent=2));return d def main(): ap=argparse.ArgumentParser();ap.add_argument('--task',choices=FILES);ap.add_argument('--method',choices=['fast-lewm','dinowm','pldm','lejepa'],default='fast-lewm');ap.add_argument('--checkpoint');ap.add_argument('--offset',type=int,default=25);ap.add_argument('--batch',type=int,default=2);ap.add_argument('--limit',type=int,default=200);ap.add_argument('--seed',type=int,choices=range(42,47),default=42);ap.add_argument('--trained',action='store_true');ap.add_argument('--smoke',action='store_true');ap.add_argument('--prepare',action='store_true');args=ap.parse_args() if args.prepare: for task in FILES: for offset in [25,50,75,100]: for seed in range(42,47): manifest(task,offset,seed) save_sheet();return torch.set_num_threads(2);torch.manual_seed(args.seed);np.random.seed(args.seed) sys.path.insert(0,str(ROOT/'repos'/('Fast-LeWorldModel' if args.method=='fast-lewm' else 'le-wm'))) cfg=OmegaConf.load(ROOT/'repos/le-wm/config/eval'/f'{args.task}.yaml') pairs=manifest(args.task,args.offset,args.seed) ckpt=Path(args.checkpoint);sha=hashlib.sha256(ckpt.read_bytes()).hexdigest() model=None if args.method=='fast-lewm':model=torch.load(ckpt,map_location='cpu',weights_only=False).eval().cuda().requires_grad_(False) elif not args.trained and args.task!='pusht':raise ValueError('The released original DINO-WM adapter currently supports PushT only') if args.method=='fast-lewm': model.consistency_loss_weight=0.; model.action_num_blocks_per_step=None plan=swm.PlanConfig(horizon=1,receding_horizon=1,action_block=25) else: plan=swm.PlanConfig(**OmegaConf.to_container(cfg.plan_config)) ds=swm.data.HDF5Dataset(path=pairs['dataset'],keys_to_cache=list(cfg.dataset.keys_to_cache)) process={} for k in cfg.dataset.keys_to_cache: x=ds.get_col_data(k);scaler=StandardScaler().fit(x[~np.isnan(x).any(axis=1)]);process[k]=scaler if k!='action':process['goal_'+k]=scaler if args.trained: from train_models import TrainedCost payload=torch.load(ckpt,map_location='cpu',weights_only=False) metadata=payload['metadata'] if not args.smoke and metadata.get('epochs_completed')!=10:raise ValueError('Only completed 10-epoch checkpoints may enter the report') if metadata['method']!=args.method or metadata['task']!=args.task:raise ValueError('Training checkpoint identity mismatch') model=TrainedCost(payload['model'],args.method).eval().cuda().requires_grad_(False) elif args.method=='dinowm': from dino_adapter import DinoAdapter model=DinoAdapter(ckpt,process).eval().requires_grad_(False) mean,std=([.5]*3,[.5]*3) if args.method=='dinowm' and not args.trained else ([.485,.456,.406],[.229,.224,.225]) transform=T.Compose([T.ToImage(),T.ToDtype(torch.float32,scale=True),T.Normalize(mean=mean,std=std),T.Resize((224,224))]) out=result_path(args.method,args.task,args.offset,args.seed) if args.smoke:out=ROOT/'manifests'/f'smoke_{args.method}_{args.task}.json' result=json.loads(out.read_text()) if out.exists() else {'method':args.method,'task':args.task,'offset':args.offset,'seed':args.seed,'successes':[],'checkpoint_sha256':sha,'eval_budget':50,'cem':{'samples':300,'iterations':30,'topk':30},'elapsed_seconds':0} if args.trained: result['checkpoint_release']=f'Locally trained {metadata["epochs_completed"]} epochs: {metadata["variant"]}' result['training']=metadata result['smoke']=args.smoke elif args.method=='dinowm': result['checkpoint_release']='original DINO-WM OSF outputs/pusht; trained on pusht_noise' result['adapter']='dino_adapter.py; original preprocessing and proprioceptive objective; common CEM coordinates' if result.get('seed',42)!=args.seed:raise ValueError('Seed mismatch') result['seed']=args.seed if result['checkpoint_sha256']!=sha:raise ValueError('Checkpoint changed during resumed evaluation') for start in range(len(result['successes']),args.limit,args.batch): end=min(start+args.batch,args.limit);t=time.time() wc=OmegaConf.to_container(cfg.world);wc['num_envs']=end-start;wc['max_episode_steps']=100 world=swm.World(**wc,image_shape=(224,224)) solver=swm.solver.CEMSolver(model,batch_size=1,num_samples=300,n_steps=30,topk=30,var_scale=1,device='cuda',seed=args.seed+start) policy=swm.policy.WorldModelPolicy(solver=solver,config=plan,process=process,transform={'pixels':transform,'goal':transform}) world.set_policy(policy) metrics=world.evaluate(dataset=ds,episodes_idx=pairs['episodes'][start:end],start_steps=pairs['start_steps'][start:end],goal_offset=args.offset,eval_budget=50,callables=OmegaConf.to_container(cfg.eval.callables),video=None) result['successes'].extend(np.asarray(metrics['episode_successes'],dtype=bool).tolist());result['elapsed_seconds']+=time.time()-t tmp=out.with_suffix('.tmp');tmp.write_text(json.dumps(result,indent=2));tmp.replace(out) if not args.smoke:save_sheet() print(f'PROGRESS {args.method} {args.task} offset={args.offset} seed={args.seed} {end}/200 successes={sum(result["successes"])} seconds={result["elapsed_seconds"]:.1f}',flush=True) world.envs.close();del world,policy,solver torch.cuda.empty_cache() if __name__=='__main__':main()