File size: 6,692 Bytes
9375a59
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""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()