transfer / actaug /code /rollout.py
Ronaldo-GOAT's picture
actaug: simulator action-augmentation generation code, Table-2 eval protocol, raw dumps of the 256 prio episodes
bd08baf verified
Raw History Blame Contribute Delete
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()