File size: 12,326 Bytes
bd08baf | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 | #!/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()
|