| """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) |
| |
| 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() |
|
|