Leplanner / code /scripts /eval_controller.py
nottygian's picture
Push scripts
dc9f917 verified
Raw
History Blame Contribute Delete
10.1 kB
"""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 # 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[1]))
from lejepa_control.solver import ControllerSolver, load_controller # noqa: E402
from lejepa_control.world_model import load_lewm # noqa: E402
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)
# execute one block then replan (spec section 9); set equal to --horizon
# for the "execute the full plan" ablation
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'
)
# world model + env to evaluate against; the PushT defaults keep existing
# runs unchanged. TwoRoom uses quentinll/lewm-tworooms + swm/TwoRoom-v1.
p.add_argument('--wm-name', default='quentinll/lewm-pusht')
p.add_argument('--env-name', default='swm/PushT-v1')
# the TwoRoom h5 stores the agent position as `pos_agent` and the episode
# id as `ep_idx`; PushT uses `state` / `episode_idx`. These map the harness
# onto either dataset without changing the env-facing info keys.
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'):
# fit on the dataset column (state_col), but store under the env
# info-key the policy matches against (always 'state'); TwoRoom's
# h5 names this column `pos_agent` but its env still emits `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}'
# the central claim is cost, so count predictor forwards rather than
# asserting them: one call = one batched latent transition
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
# terminal latent distance at each replan -> CEM regret, comparable across
# planners because both are scored under the same frozen model. GoalMSE
# sums over the latent dim while the controller averages, so rescale to
# per-dim. Done by retyping the instance: __call__ is looked up on the
# type, and the policy isinstance-checks against the Solver protocol.
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)
# held-out start/goal pairs, identical across planners for a fair compare
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
),
# Averaging over solver calls is not comparable across execution
# lengths: a run that replans every block makes more calls, and the
# later ones are taken nearer the goal. The first call is taken from
# the same held-out state by every planner, so that one is.
'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,
# per-episode outcomes: 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()
],
}
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()