#!/usr/bin/env python3 """rollout.py -- perturbed-object rollouts with GT-approach -> GR00T-grasp handoff. For one episode of a robocasa LeRobot export: 1. reset the sim to the recorded initial state (model.xml + mujoco state), 2. apply a yaw/translation perturbation to the manipulated object, 3. replay the recorded actions until just before the recorded grasp, 4. hand over to the GR00T policy for the grasp (N seeded samples per perturbation), resume the recorded transport once the object is lifted, 5. keep only task-successful attempts (actions + states + 3-cam mp4s, the SAME cameras as the source dataset). Modes: gt pure recorded-action replay (sanity ceiling; no policy server) hybrid GT approach -> policy grasp -> GT transport (unconstrained) policy GR00T end-to-end from the (perturbed) initial state constrained joints 1-6 + base pinned to the recording; the policy owns only the wrist (joint 7 servo around a searched dj7_init offset) and the gripper -- the sim_action_aug algorithm (see constrained.py) All CLI options can instead live in a YAML file passed as --config; explicit CLI flags override the YAML. The `constrained:` block of the YAML configures the constrained mode (dj7_init_deg grid, gripper_mode, caps, ...). Examples (run with the mimicgen python, MUJOCO_GL=egl): # sanity: unperturbed GT replay must succeed rollout.py --episode 0 --mode gt # 3 random perturbations x 5 policy samples each, keep successes rollout.py --episode 0 --mode hybrid --n_perturbs 3 --n_samples 5 \ --rot_deg 60 --trans_m 0.05 --port 8801 # everything from a config file rollout.py --config configs/constrained.yaml --episode 0 """ import argparse import json import time from pathlib import Path import numpy as np from actaug_core import (DEFAULT_DATASET, LIFT_M, CAMS, PolicyClient, apply_object_perturbation, env_to_parquet_action, grasp_frame, groot_obs, groot_state_vec, is_success, load_episode, make_env, obj_z, policy_chunk_to_env, render3, reset_episode, sample_perturbation, settle, to_env_action, write_mp4) def run_attempt(env, ep, mode, pert, client, attempt_seed, pre_buffer, policy_budget, resume_offset, exec_h, record): """One rollout. Returns (row, dump). dump has parquet-order executed actions, 16-d states, and 3-cam frames when `record`.""" reset_episode(env, ep) if pert is not None: dyaw, dx, dy = pert apply_object_perturbation(env, dyaw, dx, dy) settle(env) acts = ep["actions"] grasp_t = grasp_frame(acts) init_z = obj_z(env) dump = {"actions": [], "states": [], "frames": []} st = {"succ": False, "peak_lift": 0.0, "policy_steps": 0} def step(a_env): if record: dump["states"].append(groot_state_vec(env)) dump["actions"].append(env_to_parquet_action(np.asarray(a_env, float))) dump["frames"].append(render3(env)) env.step(a_env) st["peak_lift"] = max(st["peak_lift"], obj_z(env) - init_z) if is_success(env): st["succ"] = True def run_policy(budget, stop_on_lift): used, chunk_i = 0, 0 while used < budget and not st["succ"]: seed = attempt_seed * 100003 + chunk_i chunk = client.get_action(groot_obs(env, ep["instruction"], action_seed=seed)) chunk_i += 1 for j in range(exec_h): step(policy_chunk_to_env(chunk, j)) used += 1 st["policy_steps"] += 1 if st["succ"] or used >= budget or \ (stop_on_lift and obj_z(env) - init_z > LIFT_M): break if stop_on_lift and obj_z(env) - init_z > LIFT_M: break if mode == "gt": for t in range(len(acts)): step(to_env_action(acts[t])) elif mode == "policy": run_policy(len(acts) + 50, stop_on_lift=False) elif mode == "hybrid": for t in range(max(0, grasp_t - pre_buffer)): step(to_env_action(acts[t])) if not st["succ"]: run_policy(policy_budget, stop_on_lift=True) for t in range(min(len(acts) - 1, grasp_t + resume_offset), len(acts)): step(to_env_action(acts[t])) else: raise ValueError(mode) row = dict(mode=mode, attempt_seed=attempt_seed, grasp_t=grasp_t, dyaw_deg=pert[0] if pert else 0.0, dx=pert[1] if pert else 0.0, dy=pert[2] if pert else 0.0, policy_steps=st["policy_steps"], peak_lift_m=round(st["peak_lift"], 4), grasped=int(st["peak_lift"] > LIFT_M), success=int(st["succ"]), n_steps=len(dump["actions"]) if record else -1) return row, dump def save_dump(out_dir, ep, row, dump, fps): out_dir.mkdir(parents=True, exist_ok=True) np.save(out_dir / "actions.npy", np.stack(dump["actions"]).astype(np.float64)) np.savez_compressed(out_dir / "states.npz", states=np.stack(dump["states"]).astype(np.float64)) frames = dump["frames"] # list of [left, right, wrist] for i, name in enumerate(["left", "right", "wrist"]): write_mp4(out_dir / f"{name}.mp4", [f[i] for f in frames], fps=fps) write_mp4(out_dir / "3cam.mp4", [np.concatenate(f, axis=1) for f in frames], fps=fps) if dump.get("obj"): np.savez_compressed(out_dir / "obj_track.npz", obj=np.stack(dump["obj"]).astype(np.float64)) json.dump({**row, "episode": ep["ep_idx"], "instruction": ep["instruction"], "cameras": CAMS}, open(out_dir / "meta.json", "w"), indent=2) def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--config", default="", help="YAML with defaults for any " "option below + a `constrained:` block; CLI flags override") ap.add_argument("--dataset", default=str(DEFAULT_DATASET)) ap.add_argument("--episode", type=int, default=None) ap.add_argument("--mode", choices=["gt", "hybrid", "policy", "constrained"], default="hybrid") ap.add_argument("--out", default="", help="output dir " "(default outputs/ep_)") # perturbation: explicit --perturb wins; else --n_perturbs random draws ap.add_argument("--perturb", default="", help="'dyaw_deg,dx_m,dy_m' explicit " "perturbation (empty + n_perturbs 0 = nominal)") ap.add_argument("--n_perturbs", type=int, default=0, help="number of RANDOM perturbations to draw") ap.add_argument("--rot_deg", type=float, default=60.0) ap.add_argument("--rot_min_deg", type=float, default=0.0, help="minimum |yaw| of random perturbations (sign random)") ap.add_argument("--trans_m", type=float, default=0.05) ap.add_argument("--trans_min_m", type=float, default=0.0) ap.add_argument("--seed", type=int, default=0) # policy handoff ap.add_argument("--n_samples", type=int, default=1, help="policy samples per perturbation (N)") ap.add_argument("--pre_buffer", type=int, default=10, help="hand over to the policy this many steps BEFORE the recorded grasp") ap.add_argument("--policy_budget", type=int, default=200) ap.add_argument("--resume_offset", type=int, default=0, help="resume recorded transport at grasp_t + this") ap.add_argument("--exec_h", type=int, default=16, help="actions executed per chunk") ap.add_argument("--host", default="127.0.0.1") ap.add_argument("--port", type=int, default=8801) ap.add_argument("--save", choices=["success", "all", "none"], default="success", help="which attempts get a full dump (actions/states/videos)") ap.add_argument("--fps", type=int, default=20) # YAML config: values become argparse defaults, so explicit CLI flags win. import sys cons_cfg = {} if "--config" in sys.argv: import yaml cfg_path = sys.argv[sys.argv.index("--config") + 1] cfg = yaml.safe_load(open(cfg_path)) or {} cons_cfg = cfg.pop("constrained", {}) or {} known = {a.dest for a in ap._actions} bad = set(cfg) - known if bad: ap.error(f"unknown keys in {cfg_path}: {sorted(bad)}") ap.set_defaults(**cfg) args = ap.parse_args() if args.episode is None: ap.error("--episode is required (CLI or YAML)") out = Path(args.out) if args.out else \ Path(__file__).resolve().parent / "outputs" / f"ep{args.episode:06d}_{args.mode}" out.mkdir(parents=True, exist_ok=True) ep = load_episode(args.dataset, args.episode) env = make_env(ep["env_args"]) need_policy = args.mode in ("hybrid", "policy", "constrained") client = PolicyClient(args.host, args.port) if need_policy else None rng = np.random.default_rng(args.seed) if args.perturb: perts = [tuple(float(x) for x in args.perturb.split(","))] elif args.n_perturbs > 0: perts = [sample_perturbation(rng, args.rot_deg, args.trans_m, args.trans_min_m, args.rot_min_deg) for _ in range(args.n_perturbs)] else: perts = [None] n_samples = args.n_samples if need_policy else 1 # constrained mode searches a wrist-offset grid on top of the N samples dj7_grid = [float(x) for x in cons_cfg.get("dj7_init_deg", [0.0])] \ if args.mode == "constrained" else [None] rows = [] for pi, pert in enumerate(perts): for di, dj7 in enumerate(dj7_grid): for k in range(n_samples): attempt_seed = args.seed * 1000 + pi * 100 + di * 10 + k record = args.save != "none" t0 = time.time() if args.mode == "constrained": from constrained import inject_weld, run_constrained_attempt if cons_cfg.get("rel6d", True) and "actaug_grasp_weld" \ not in ep["xml"]: ep = {**ep, "xml": inject_weld( ep["xml"], cons_cfg.get("weld_solref", "0.005 1"))} reset_episode(env, ep) if pert is not None: apply_object_perturbation(env, *pert) settle(env) row, dump = run_constrained_attempt( env, ep, client, dj7, attempt_seed, cons_cfg, record) else: row, dump = run_attempt(env, ep, args.mode, pert, client, attempt_seed, args.pre_buffer, args.policy_budget, args.resume_offset, args.exec_h, record) row.update(episode=args.episode, pert_idx=pi, sample=k, dyaw_deg=pert[0] if pert else 0.0, dx=pert[1] if pert else 0.0, dy=pert[2] if pert else 0.0, wall_s=round(time.time() - t0, 1)) rows.append(row) tag = f"pert{pi:02d}" + \ (f"_dj{dj7:+04.0f}" if dj7 is not None else "") + f"_s{k:02d}" if record and (args.save == "all" or row["success"]): save_dump(out / ("accepted" if row["success"] else "rejected") / tag, ep, row, dump, args.fps) print(f"[{tag}] pert={pert} seed={attempt_seed} " f"lift={row['peak_lift_m']:+.3f} grasped={row['grasped']} " f"success={row['success']} ({row['wall_s']}s)", flush=True) import pandas as pd pd.DataFrame(rows).to_csv(out / "results.csv", index=False) n_ok = sum(r["success"] for r in rows) print(f"DONE ep{args.episode} {args.mode}: {n_ok}/{len(rows)} successful " f"-> {out}", flush=True) if __name__ == "__main__": main()