planning-baselines / code /evaluate.py
greycat111's picture
Add checkpoint provenance, configurations, licenses and five-seed results
9375a59 verified
Raw
History Blame Contribute Delete
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()