transfer / actaug /code /actaug_core.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
11.7 kB
#!/usr/bin/env python3
"""actaug_core -- self-contained helpers for object-perturbed RoboCasa rollouts.
Handoff module: imports ONLY installed libraries (robocasa/robosuite/robomimic/
numpy/pandas/zmq/torch/imageio). No imports from any other project directory.
Run with the sim interpreter:
/lp-dev/jonghoon/mimicgen_augment/envs/mimicgen/bin/python
and MUJOCO_GL=egl MUJOCO_EGL_DEVICE_ID=<gpu>.
Covers:
* loading one episode of a robocasa LeRobot export (model.xml.gz + states.npz +
ep_meta.json under extras/, 12-d actions from data/chunk-*/episode_*.parquet)
* building the env from the export's env_args and resetting to the episode
* perturbing the manipulated object's initial pose (world-z yaw + xy translation)
* GR00T-N1.5 policy client (ZMQ torch.save protocol) + obs/action converters
"""
import gzip
import io
import json
from pathlib import Path
import numpy as np
DEFAULT_DATASET = Path(
"/lp-dev/jonghoon/robocasa_full/pickplace_target_human/PickPlaceCounterToCabinet")
CAMS = ["robot0_agentview_left", "robot0_agentview_right", "robot0_eye_in_hand"]
VIEW_KEYS = ["video.left_view", "video.right_view", "video.wrist_view"]
IMG = 256
LIFT_M = 0.03 # object raised this much above its start height counts as lifted
OBJ_JOINT = "obj_joint0"
ACTION_SEED_KEY = "_policy_action_seed" # honored by myGR00T policy.get_action
_ENV_CACHE = {}
# ---------------------------------------------------------------- episode I/O
def load_episode(dataset_root, ep_idx):
"""Read everything needed to re-simulate episode `ep_idx` of a robocasa
LeRobot export. Returns a dict; heavy arrays are numpy."""
root = Path(dataset_root)
ed = root / "extras" / f"episode_{ep_idx:06d}"
xml = gzip.open(ed / "model.xml.gz", "rt").read()
states = np.load(ed / "states.npz")["states"]
ep_meta = json.load(open(ed / "ep_meta.json"))
import pandas as pd
chunk = ep_idx // 1000
pq = root / "data" / f"chunk-{chunk:03d}" / f"episode_{ep_idx:06d}.parquet"
df = pd.read_parquet(pq)
actions = np.stack(df["action"].to_numpy()).astype(np.float64) # (T, 12) parquet order
# instruction: tasks.jsonl row addressed by the annotation column (fallback: ep_meta lang)
instr = ep_meta.get("lang", "")
try:
tasks = {}
for line in open(root / "meta" / "tasks.jsonl"):
row = json.loads(line)
tasks[int(row["task_index"])] = row["task"]
instr = tasks[int(df["annotation.human.task_description"].iloc[0])]
except Exception:
pass
env_args = json.load(open(root / "extras" / "dataset_meta.json"))["env_args"]
return dict(ep_idx=ep_idx, xml=xml, states=states, ep_meta=ep_meta,
actions=actions, instruction=instr, env_args=env_args)
def make_env(env_args):
"""EnvRobosuite for the export's env_args (cached per env_name)."""
name = env_args["env_name"]
if name in _ENV_CACHE:
return _ENV_CACHE[name]
import robocasa # noqa: F401 (registers kitchen envs)
import robomimic.utils.obs_utils as ObsUtils
ObsUtils.initialize_obs_utils_with_obs_specs({"obs": {"low_dim": [], "rgb": []}})
from robomimic.envs.env_robosuite import EnvRobosuite
kw = dict(env_args["env_kwargs"])
kw.pop("env_name", None)
env = EnvRobosuite(name, render=False, render_offscreen=True, use_image_obs=False,
camera_names=CAMS, camera_heights=IMG, camera_widths=IMG, **kw)
_ENV_CACHE[name] = env
return env
def reset_episode(env, ep):
"""Reset env to the recorded initial sim state of `ep` (from load_episode)."""
env.reset_to({"model": ep["xml"], "states": ep["states"][0],
"ep_meta": json.dumps(ep["ep_meta"])})
# ------------------------------------------------------------- perturbation
def _q_mul(a, b): # wxyz
w0, x0, y0, z0 = a
w1, x1, y1, z1 = b
return np.array([w0*w1 - x0*x1 - y0*y1 - z0*z1,
w0*x1 + x0*w1 + y0*z1 - z0*y1,
w0*y1 - x0*z1 + y0*w1 + z0*x1,
w0*z1 + x0*y1 - y0*x1 + z0*w1])
def apply_object_perturbation(env, dyaw_deg=0.0, dx=0.0, dy=0.0, dz=0.0):
"""Spin the manipulated object about the WORLD vertical and translate it in
the world frame, in-sim, right after reset. Returns (pos_before, pos_after)."""
sim = env.env.sim
jid = sim.model.joint_name2id(OBJ_JOINT)
adr = int(sim.model.jnt_qposadr[jid])
p0 = sim.data.qpos[adr:adr + 3].copy()
th = np.radians(dyaw_deg) / 2.0
qz = np.array([np.cos(th), 0.0, 0.0, np.sin(th)])
q = _q_mul(qz, sim.data.qpos[adr + 3:adr + 7].copy())
sim.data.qpos[adr + 3:adr + 7] = q / np.linalg.norm(q)
sim.data.qpos[adr:adr + 3] = p0 + np.array([dx, dy, dz])
sim.forward()
return p0, sim.data.qpos[adr:adr + 3].copy()
def sample_perturbation(rng, rot_deg=180.0, trans_m=0.05, trans_min_m=0.0,
rot_min_deg=0.0):
"""Yaw magnitude uniform in [rot_min_deg, rot_deg] with random sign;
translation direction uniform on the circle, radius uniform in
[trans_min_m, trans_m]. Returns (dyaw, dx, dy)."""
dyaw = float(rng.uniform(rot_min_deg, rot_deg)) * (1 if rng.random() < 0.5 else -1)
r = float(rng.uniform(trans_min_m, trans_m))
phi = float(rng.uniform(0, 2 * np.pi))
return dyaw, r * np.cos(phi), r * np.sin(phi)
def settle(env, n_steps=10):
"""Let physics settle after a perturbation (object may be slightly off the
counter surface after a yaw about its joint origin)."""
sim = env.env.sim
for _ in range(n_steps):
sim.step()
sim.forward()
# ------------------------------------------------------ actions & obs bridges
def to_env_action(a):
"""parquet 12-d (base_motion[0:4], control_mode[4], ee_pos[5:8], ee_rot[8:11],
grip[11]) -> robosuite env-order 12-d."""
e = np.zeros(12)
e[0:3] = a[5:8]
e[3:6] = a[8:11]
e[6] = 1.0 if a[11] > 0 else -1.0
e[7:11] = a[0:4]
e[11] = 1.0 if a[4] > 0 else -1.0
return e
def env_to_parquet_action(e):
"""Inverse of to_env_action (for dumping executed trajectories)."""
a = np.zeros(12)
a[0:4] = e[7:11]
a[4] = 1.0 if e[11] > 0 else -1.0
a[5:8] = e[0:3]
a[8:11] = e[3:6]
a[11] = 1.0 if e[6] > 0 else -1.0
return a
def policy_chunk_to_env(chunk, j):
"""GR00T action-chunk step j -> env-order 12-d (binarized gripper/control_mode)."""
eep = np.asarray(chunk["action.end_effector_position"])[j].ravel()[:3]
eer = np.asarray(chunk["action.end_effector_rotation"])[j].ravel()[:3]
grip = float(np.asarray(chunk["action.gripper_close"])[j].ravel()[0])
base = np.asarray(chunk["action.base_motion"])[j].ravel()[:4]
ctrl = float(np.asarray(chunk["action.control_mode"])[j].ravel()[0])
e = np.zeros(12)
e[0:3] = eep
e[3:6] = eer
e[6] = 1.0 if grip > 0 else -1.0
e[7:11] = base
e[11] = 1.0 if ctrl > 0 else -1.0
return e
def grasp_frame(actions):
"""First control step whose recorded gripper bit commands CLOSE."""
close = np.where(actions[:, 11] > 0)[0]
return int(close[0]) if len(close) else len(actions) // 2
def render3(env):
sim = env.env.sim
return [sim.render(camera_name=c, width=IMG, height=IMG)[::-1] for c in CAMS]
def groot_state_vec(env):
"""16-d observation.state in the export's modality.json layout."""
di = env.env._get_observations(force_update=False)
return np.concatenate([
np.asarray(di["robot0_base_pos"]).ravel()[:3],
np.asarray(di["robot0_base_quat"]).ravel()[:4],
np.asarray(di["robot0_base_to_eef_pos"]).ravel()[:3],
np.asarray(di["robot0_base_to_eef_quat"]).ravel()[:4],
np.asarray(di["robot0_gripper_qpos"]).ravel()[:2],
]).astype(np.float64)
def groot_obs(env, instruction, action_seed=None):
di = env.env._get_observations(force_update=True)
l, r, w = render3(env)
obs = {"video.left_view": l[None], "video.right_view": r[None],
"video.wrist_view": w[None],
"state.end_effector_position_relative": np.asarray(di["robot0_base_to_eef_pos"])[None],
"state.end_effector_rotation_relative": np.asarray(di["robot0_base_to_eef_quat"])[None],
"state.gripper_qpos": np.asarray(di["robot0_gripper_qpos"])[None],
"state.base_position": np.asarray(di["robot0_base_pos"])[None],
"state.base_rotation": np.asarray(di["robot0_base_quat"])[None],
"annotation.human.action.task_description": [instruction]}
if action_seed is not None:
obs[ACTION_SEED_KEY] = int(action_seed)
return obs
def groot_obs_oxe(env, instruction, action_seed=None):
"""DROID/oxe_droid-format obs for the BASE (non-finetuned) GR00T-N1.5, whose
only viable manipulation head here is oxe_droid (best-effort frame mapping,
same approach as gripper_state_replay.groot_obs_oxe)."""
import robosuite.utils.transform_utils as T
di = env.env._get_observations(force_update=True)
l, r, w = render3(env)
quat = np.asarray(di["robot0_base_to_eef_quat"]) # [x,y,z,w]
euler = T.mat2euler(T.quat2mat(quat))
gq = np.asarray(di["robot0_gripper_qpos"])
g01 = float(np.clip(np.mean(np.abs(gq)) / 0.04, 0.0, 1.0)) # 1=open, 0=closed
obs = {"video.exterior_image_1": l[None], "video.exterior_image_2": r[None],
"video.wrist_image": w[None],
"state.eef_position": np.asarray(di["robot0_base_to_eef_pos"])[None],
"state.eef_rotation": np.asarray(euler)[None],
"state.gripper_position": np.array([[g01]]),
"annotation.language.language_instruction": [instruction]}
if action_seed is not None:
obs[ACTION_SEED_KEY] = int(action_seed)
return obs
def oxe_chunk_to_env(chunk, j, grip_flip=False):
"""oxe_droid action step -> robocasa env-order 12-d (base=0, control_mode=-1).
DROID gripper_position ~[0,1]; default close = g > 0.5 (flip to invert)."""
dp = np.asarray(chunk["action.eef_position_delta"])[j].ravel()[:3]
dr = np.asarray(chunk["action.eef_rotation_delta"])[j].ravel()[:3]
g = float(np.asarray(chunk["action.gripper_position"])[j].ravel()[0])
e = np.zeros(12)
e[0:3] = dp
e[3:6] = dr
close = (g < 0.5) if grip_flip else (g > 0.5)
e[6] = 1.0 if close else -1.0
e[11] = -1.0
return e
def obj_z(env):
return float(env.env._get_observations(force_update=False)["obj_pos"][2])
def is_success(env):
return bool(env.is_success().get("task", False))
# ------------------------------------------------------------- policy client
class PolicyClient:
"""ZMQ REQ client for the myGR00T RobotInferenceServer (torch.save protocol)."""
def __init__(self, host="127.0.0.1", port=8801):
import zmq
self.ctx = zmq.Context()
self.sock = self.ctx.socket(zmq.REQ)
self.sock.connect(f"tcp://{host}:{port}")
def get_action(self, obs):
import torch
buf = io.BytesIO()
torch.save({"endpoint": "get_action", "data": obs}, buf)
self.sock.send(buf.getvalue())
msg = self.sock.recv()
if msg == b"ERROR":
raise RuntimeError("policy server returned ERROR")
return torch.load(io.BytesIO(msg), weights_only=False)
# ------------------------------------------------------------------- output
def write_mp4(path, frames, fps=20):
import imageio
w = imageio.get_writer(str(path), fps=fps, codec="libx264",
macro_block_size=1, ffmpeg_params=["-pix_fmt", "yuv420p"])
for f in frames:
w.append_data(np.asarray(f).astype(np.uint8))
w.close()