Download code/scripts/eval_planner.py from SaltedLemon/lejepa-control-pusht: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/SaltedLemon/lejepa-control-pusht/resolve/main/code/scripts/eval_planner.py
- Command line
-
hf download hf://SaltedLemon/lejepa-control-pusht/code/scripts/eval_planner.py
-
curl -L -o eval_planner.py https://huggingface.co/SaltedLemon/lejepa-control-pusht/resolve/main/code/scripts/eval_planner.py
11.2 kB
| """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]) | |
| def action_dim(self): | |
| return self._action_dim * self._action_block | |
| def n_envs(self): | |
| return self._n_envs | |
| 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() | |