Download code/evaluate.py from FidelityWM/planning-baselines: direct link, hf CLI and curl.
- Browser
- Download file 6.69 kB
-
https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/evaluate.py
- Command line
-
hf download hf://FidelityWM/planning-baselines/code/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/evaluate.py
6.69 kB
| """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() | |