lejepa-control-pusht / code /scripts /eval_planner.py
SaltedLemon's picture
Upload code/scripts/eval_planner.py with huggingface_hub
66c77b9 verified
Raw
History Blame Contribute Delete
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])
@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()