| """Evaluate the recursive planner in the real PushT simulator. |
| |
| Mirrors ``scripts/eval_controller.py`` exactly — same env, same wrappers, same |
| preprocessing, same held-out start/goal pairs drawn from the same seed — so |
| that planner, baseline controller, CEM and a random-action floor are all |
| scored on identical episodes. Every gate in the staged bring-up is decided on |
| real-env success, never on the world model's own imagined distance. |
| |
| ``--planner random`` is the stage-A gate's floor. ``--planner controller`` is |
| the phase-1 baseline the recursion has to beat. |
| """ |
|
|
| 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[2])) |
|
|
| from lejepa_control.solver import ControllerSolver, load_controller |
| from lejepa_control.world_model import load_lewm |
| from lejepa_control_2.solver import PlannerSolver, load_planner |
|
|
|
|
| class RandomSolver: |
| """Uniform action blocks — the floor the stage-A gate is measured against. |
| |
| A planner that does not beat this is not planning, whatever its training |
| loss is doing. |
| """ |
|
|
| def __init__(self, horizon=5, seed=0): |
| self._horizon = horizon |
| self._n_envs = 1 |
| self._action_dim = 2 |
| self._action_block = 5 |
| self._gen = torch.Generator().manual_seed(seed) |
|
|
| def configure(self, *, action_space, n_envs, config): |
| self._n_envs = n_envs |
| self._horizon = config.horizon |
| self._action_block = config.action_block |
| self._action_dim = int(action_space.shape[-1]) |
|
|
| @property |
| def action_dim(self): |
| return self._action_dim * self._action_block |
|
|
| @property |
| def n_envs(self): |
| return self._n_envs |
|
|
| @property |
| def horizon(self): |
| return self._horizon |
|
|
| def solve(self, info_dict, init_action=None): |
| b = info_dict['pixels'].shape[0] |
| actions = torch.rand( |
| b, self._horizon, self.action_dim, generator=self._gen |
| ) * 2 - 1 |
| return {'actions': actions, 'costs': torch.zeros(b)} |
|
|
| __call__ = solve |
|
|
|
|
| def parse_args(): |
| p = argparse.ArgumentParser() |
| p.add_argument( |
| '--planner', |
| default='planner', |
| choices=['planner', 'controller', 'cem', 'random'], |
| ) |
| p.add_argument('--checkpoint', default='data/runs/planner/planner.pt') |
| p.add_argument( |
| '--controller', default='data/runs/controller/controller.pt' |
| ) |
| p.add_argument('--cycles', type=int, default=None) |
| p.add_argument('--inner', type=int, default=None) |
| 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('--out', default='data/runs/eval_planner') |
| p.add_argument('--tag', default=None) |
| |
| |
| 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 build_solver(args, model, device, latent_dim): |
| """Returns ``(solver, tag, extra)``.""" |
| if args.planner == 'planner': |
| planner, ckpt = load_planner( |
| args.checkpoint, |
| device=device, |
| cycles=args.cycles, |
| inner=args.inner, |
| horizon=args.horizon, |
| ) |
| solver = PlannerSolver( |
| model, planner, device=device, |
| cycles=args.cycles, inner=args.inner, |
| ) |
| stage = ckpt['args'].get('stage') or 'custom' |
| tag = f'planner_{stage}_T{planner.cycles}_n{planner.inner}' |
| print( |
| f'planner from step {ckpt["step"]}, stage {stage}, ' |
| f'T={planner.cycles} n={planner.inner} H={planner.horizon}' |
| ) |
| return solver, tag, {'step': ckpt['step'], 'stage': stage} |
|
|
| 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) |
| print(f'baseline controller from step {ckpt["step"]}') |
| return ( |
| solver, |
| f'controller_K{controller.refinements}', |
| {'step': ckpt['step']}, |
| ) |
|
|
| if args.planner == 'random': |
| return RandomSolver(args.horizon, args.seed), 'random', {} |
|
|
| 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, |
| ) |
| return solver, f'cem_s{args.cem_samples}_n{args.cem_steps}', {} |
|
|
|
|
| def main(): |
| args = parse_args() |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' |
|
|
| world = swm.World( |
| env_name='swm/PushT-v1', |
| 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', 'state'], |
| ) |
|
|
| process = {} |
| for col in ('action', 'proprio', 'state'): |
| data = dataset.get_col_data(col) |
| 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(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, |
| ) |
|
|
| solver, tag, extra = build_solver(args, model, device, latent_dim) |
| tag = args.tag or tag |
|
|
| 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, per_cycle_trace = [], [] |
| cost_scale = 1.0 / latent_dim if args.planner == 'cem' else 1.0 |
| 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: |
| terminals.append( |
| float(torch.as_tensor(costs).float().mean()) * cost_scale |
| ) |
| if out.get('per_cycle') is not None: |
| per_cycle_trace.append(out['per_cycle']) |
| 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('episode_idx') |
| 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': 'state'}}}, |
| { |
| 'method': '_set_goal_state', |
| 'args': {'goal_state': {'value': 'goal_state'}}, |
| }, |
| ], |
| video=out_dir if args.video else None, |
| ) |
| elapsed = time.time() - t0 |
|
|
| result = { |
| 'planner': tag, |
| 'kind': args.planner, |
| '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, |
| '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, |
| 'seed': args.seed, |
| 'cycles': args.cycles, |
| 'inner': args.inner, |
| |
| |
| 'episode_successes': [ |
| bool(x) for x in metrics['episode_successes'].tolist() |
| ], |
| **extra, |
| } |
| if per_cycle_trace: |
| result['mean_per_cycle'] = ( |
| np.asarray(per_cycle_trace).mean(axis=0).tolist() |
| ) |
|
|
| print(json.dumps( |
| {k: v for k, v in result.items() if k != 'episode_successes'}, indent=2 |
| )) |
| with (out_dir / 'results.jsonl').open('a') as f: |
| f.write(json.dumps(result) + '\n') |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|