actaug: simulator action-augmentation generation code, Table-2 eval protocol, raw dumps of the 256 prio episodes
bd08baf verified Download actaug/code/rollout.py from Ronaldo-GOAT/transfer: direct link, hf CLI and curl.
- Browser
- Download file 12.3 kB
-
https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/actaug/code/rollout.py
- Command line
-
hf download hf://Ronaldo-GOAT/transfer/actaug/code/rollout.py
-
curl -L -o rollout.py https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/actaug/code/rollout.py
12.3 kB
| #!/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<EP>_<mode>)") | |
| # 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() | |