"""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 # noqa: E402 import json # noqa: E402 import sys # noqa: E402 import time # noqa: E402 from pathlib import Path # noqa: E402 import hdf5plugin # noqa: F401,E402 -- blosc filter for the expert h5 import numpy as np # noqa: E402 import stable_pretraining as spt # noqa: E402 import stable_worldmodel as swm # noqa: E402 import torch # noqa: E402 from sklearn import preprocessing # noqa: E402 from torchvision.transforms import v2 as transforms # noqa: E402 sys.path.insert(0, str(Path(__file__).resolve().parents[2])) from lejepa_control.solver import ControllerSolver, load_controller # noqa: E402 from lejepa_control.world_model import load_lewm # noqa: E402 from lejepa_control_2.solver import PlannerSolver, load_planner # noqa: E402 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) # the same seed must be used for every configuration so the held-out # start/goal pairs are identical and the comparison stays paired 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 # GoalMSE sums over the latent dim while the planner averages, so CEM's # cost is rescaled to per-dim to stay comparable 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) # held-out start/goal pairs, identical across planners for a fair compare 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 ), # the first call is taken from the same held-out state by every # planner, so unlike the mean it is comparable across execution lengths '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, # all rows share start/goal pairs, so planner comparisons must be # paired rather than treated as independent '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()