| """Evaluate the controller in the real PushT simulator, against CEM.
|
|
|
| Both planners run through the same ``WorldModelPolicy``, the same wrappers and
|
| the same held-out initial-state/goal pairs, so the only thing that differs is
|
| how the plan is produced. Reports success rate and wall-clock planning time —
|
| the central claim is approaching CEM's success with far fewer world-model
|
| evaluations.
|
| """
|
|
|
| import os
|
|
|
| os.environ['MUJOCO_GL'] = 'egl'
|
|
|
| import argparse
|
| import json
|
| import sys
|
| import time
|
| from pathlib import Path
|
|
|
| import hdf5plugin
|
| import numpy as np
|
| import stable_pretraining as spt
|
| import stable_worldmodel as swm
|
| import torch
|
| from sklearn import preprocessing
|
| from torchvision.transforms import v2 as transforms
|
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
|
|
| from lejepa_control.solver import ControllerSolver, load_controller
|
| from lejepa_control.world_model import load_lewm
|
|
|
|
|
| def parse_args():
|
| p = argparse.ArgumentParser()
|
| p.add_argument('--controller', default='data/runs/controller/controller.pt')
|
| p.add_argument('--planner', default='controller', choices=['controller', 'cem'])
|
| p.add_argument('--refinements', type=int, default=None)
|
| p.add_argument('--num-eval', type=int, default=50)
|
| p.add_argument('--eval-budget', type=int, default=50)
|
| p.add_argument('--goal-offset', type=int, default=25)
|
| p.add_argument('--horizon', type=int, default=5)
|
|
|
|
|
| p.add_argument('--receding-horizon', type=int, default=1)
|
| p.add_argument('--cem-samples', type=int, default=300)
|
| p.add_argument('--cem-steps', type=int, default=30)
|
| p.add_argument(
|
| '--dataset', default='data/swm_home/datasets/pusht_expert_train.h5'
|
| )
|
|
|
|
|
| p.add_argument('--wm-name', default='quentinll/lewm-pusht')
|
| p.add_argument('--env-name', default='swm/PushT-v1')
|
|
|
|
|
|
|
| p.add_argument('--state-col', default='state')
|
| p.add_argument('--episode-col', default='episode_idx')
|
| p.add_argument('--out', default='data/runs/eval')
|
| p.add_argument('--seed', type=int, default=42)
|
| p.add_argument('--video', action='store_true')
|
| return p.parse_args()
|
|
|
|
|
| def img_transform(size=224):
|
| return transforms.Compose(
|
| [
|
| transforms.ToImage(),
|
| transforms.ToDtype(torch.float32, scale=True),
|
| transforms.Normalize(**spt.data.dataset_stats.ImageNet),
|
| transforms.Resize(size=size),
|
| ]
|
| )
|
|
|
|
|
| def main():
|
| args = parse_args()
|
| device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
|
|
| world = swm.World(
|
| env_name=args.env_name,
|
| num_envs=args.num_eval,
|
| max_episode_steps=2 * args.eval_budget,
|
| image_shape=(224, 224),
|
| )
|
|
|
| dataset = swm.data.load_dataset(
|
| str(Path(args.dataset).resolve()),
|
| keys_to_cache=['action', 'proprio', args.state_col],
|
| )
|
|
|
| process = {}
|
| for col in ('action', 'proprio', 'state'):
|
|
|
|
|
|
|
| src = args.state_col if col == 'state' else col
|
| data = dataset.get_col_data(src)
|
| data = data[~np.isnan(data).any(axis=1)]
|
| scaler = preprocessing.StandardScaler().fit(data)
|
| process[col] = scaler
|
| if col != 'action':
|
| process[f'goal_{col}'] = scaler
|
|
|
| transform = {'pixels': img_transform(), 'goal': img_transform()}
|
|
|
| model = load_lewm(name=args.wm_name, device=device)
|
| model.interpolate_pos_encoding = True
|
| latent_dim = model.predictor.input_dim
|
|
|
| config = swm.PlanConfig(
|
| horizon=args.horizon,
|
| receding_horizon=args.receding_horizon,
|
| action_block=5,
|
| history_len=model.predictor.num_frames,
|
| )
|
|
|
| if args.planner == 'controller':
|
| controller, ckpt = load_controller(
|
| args.controller,
|
| latent_dim=latent_dim,
|
| device=device,
|
| refinements=args.refinements,
|
| )
|
| solver = ControllerSolver(model, controller, device=device)
|
| tag = f'controller_K{controller.refinements}'
|
| print(f'controller from step {ckpt["step"]}, K={controller.refinements}')
|
| else:
|
| cost = swm.planning.ShootingCostEvaluator(model, swm.planning.GoalMSE())
|
| solver = swm.planning.CEMSolver(
|
| cost=cost,
|
| num_samples=args.cem_samples,
|
| n_steps=args.cem_steps,
|
| topk=30,
|
| device=device,
|
| )
|
| tag = f'cem_s{args.cem_samples}_n{args.cem_steps}'
|
|
|
|
|
|
|
| calls = {'n': 0, 'rows': 0}
|
| inner_predict = model.predictor.forward
|
|
|
| def counting_predict(*a, **kw):
|
| calls['n'] += 1
|
| first = a[0] if a else next(iter(kw.values()))
|
| calls['rows'] += first.shape[0]
|
| return inner_predict(*a, **kw)
|
|
|
| model.predictor.forward = counting_predict
|
|
|
|
|
|
|
|
|
|
|
|
|
| terminals = []
|
| scale = 1.0 if args.planner == 'controller' else 1.0 / latent_dim
|
| base = type(solver)
|
|
|
| class RecordingSolver(base):
|
| def __call__(self, info_dict, init_action=None):
|
| out = base.__call__(self, info_dict, init_action)
|
| costs = out.get('costs')
|
| if costs is not None:
|
| value = float(torch.as_tensor(costs).float().mean())
|
| terminals.append(value * scale)
|
| return out
|
|
|
| solver.__class__ = RecordingSolver
|
|
|
| policy = swm.policy.WorldModelPolicy(
|
| solver=solver,
|
| config=config,
|
| process=process,
|
| transform=transform,
|
| history_keys=('pixels',),
|
| )
|
| world.set_policy(policy)
|
|
|
|
|
| ep_idx = dataset.get_col_data(args.episode_col)
|
| step_idx = dataset.get_col_data('step_idx')
|
| episodes = np.unique(ep_idx)
|
| lengths = {e: step_idx[ep_idx == e].max() + 1 for e in episodes}
|
| max_start = np.array([lengths[e] for e in ep_idx]) - args.goal_offset - 1
|
| valid = np.nonzero(step_idx <= max_start)[0]
|
|
|
| rng = np.random.default_rng(args.seed)
|
| picked = np.sort(valid[rng.choice(len(valid), args.num_eval, replace=False)])
|
|
|
| out_dir = Path(args.out)
|
| out_dir.mkdir(parents=True, exist_ok=True)
|
|
|
| t0 = time.time()
|
| metrics = world.evaluate(
|
| dataset=dataset,
|
| start_steps=step_idx[picked].tolist(),
|
| goal_offset=args.goal_offset,
|
| eval_budget=args.eval_budget,
|
| episodes_idx=ep_idx[picked].tolist(),
|
| callables=[
|
| {
|
| 'method': '_set_state',
|
| 'args': {'state': {'value': args.state_col}},
|
| },
|
| {
|
| 'method': '_set_goal_state',
|
| 'args': {'goal_state': {'value': f'goal_{args.state_col}'}},
|
| },
|
| ],
|
| video=out_dir if args.video else None,
|
| )
|
| elapsed = time.time() - t0
|
|
|
| result = {
|
| 'planner': tag,
|
| 'env': args.env_name,
|
| 'wm': args.wm_name,
|
| 'receding_horizon': args.receding_horizon,
|
| 'success_rate': float(metrics['success_rate']),
|
| 'seconds': elapsed,
|
| 'seconds_per_episode': elapsed / args.num_eval,
|
| 'mean_terminal_distance': (
|
| float(np.mean(terminals)) if terminals else None
|
| ),
|
|
|
|
|
|
|
|
|
| 'first_terminal_distance': float(terminals[0]) if terminals else None,
|
| 'terminal_distance_trace': [float(t) for t in terminals],
|
| 'predictor_calls': calls['n'],
|
| 'predictor_rows_per_episode': calls['rows'] / args.num_eval,
|
| 'num_eval': args.num_eval,
|
| 'eval_budget': args.eval_budget,
|
| 'goal_offset': args.goal_offset,
|
| 'checkpoint': args.controller if args.planner == 'controller' else None,
|
| 'seed': args.seed,
|
| 'refinements': args.refinements,
|
|
|
|
|
| 'episode_successes': [
|
| bool(x) for x in metrics['episode_successes'].tolist()
|
| ],
|
| }
|
| print(json.dumps({k: v for k, v in result.items()
|
| if k != 'terminal_distance_trace'}, indent=2))
|
|
|
| with (out_dir / 'results.jsonl').open('a') as f:
|
| f.write(json.dumps(result) + '\n')
|
|
|
|
|
| if __name__ == '__main__':
|
| main()
|
|
|