#!/usr/bin/env python3 """constrained.py -- constrained-policy grasp re-derivation for robocasa episodes. Port of the sim_action_aug algorithm (code/rollout/r02_rollout.py, see ws/ROLLOUT_REPORT.md section 3) onto robocasa's PandaOmron, using the joint- pinning machinery proven by gripper_state_replay's grip_state/grip_policy modes: Phase 0 (approach, frames 0 .. f_pol0): arm joints 1-7 + mobile base + torso PINNED to the recorded MuJoCo state path (re-pinned every physics substep with finite-diff qvel); gripper follows the recorded open/close bit. The object's free joint is NEVER written -> fully dynamic (and carries the initial perturbation). Phase A (policy, f_pol0 = grasp_t - start_off .. grasp_t + k_post): every arm joint NOT in `policy_joints` (+ base + torso) stays pinned to the recording. Each joint j in `policy_joints` gets a policy servo: the chunk's delta EEF rotation projected onto joint j's world axis, capped `step_cap` rad per control step and `servo_caps[j]` rad in total. Joint 7 additionally carries the searched dj7_init offset (|dj7_init + servo| <= j7_total_cap) and the (-pi/2, pi/2] wrap (2-jaw symmetry about the roll axis). gripper = ALWAYS the policy's action.gripper_close: 'bit' close iff the binarized command says close (robocasa-native) 'threshold' latch full close once the raw close fraction >= grip_tau 'continuous' ctrl interpolated open->close by the raw fraction 'recorded' IGNORE the policy gripper, follow the recorded close bit (diagnostic baseline = grip_state; with all servo caps 0 no inference runs at all) Phase B (transport): everything pinned to the recording, joint 7 keeps the frozen dj7, gripper held closed until the recorded release frame, then follows the recorded opening schedule (never re-closes). +settle frames. Success = env.is_success()["task"] at any point; grasped = peak lift > 3 cm while any gripper geom touches the object (contact-gated like r02, so an object bumped upward doesn't count). """ import numpy as np from actaug_core import (LIFT_M, OBJ_JOINT, env_to_parquet_action, grasp_frame, groot_obs, groot_obs_oxe, groot_state_vec, is_success, obj_z, render3, to_env_action) DEFAULTS = dict( start_off=4, # policy takes over this many frames before grasp_t k_post=12, # policy keeps control this many frames after grasp_t replan=8, # re-query the policy every N control steps dj7_init_deg=[0.0], # searched initial wrist offsets (degrees) # WHICH arm joints the policy may servo (1-indexed Franka joints; the # gripper is ALWAYS the policy's unless gripper_mode == "recorded"). # Joints not listed stay pinned to the recording. The policy's delta EEF # rotation is projected onto each listed joint's world axis, capped # step_cap rad per control step and servo_caps[j] rad in total; joint 7 # additionally carries the searched dj7_init offset (with j7_total_cap on # dj7_init + servo) and the half-pi jaw-symmetry wrap. policy_joints=[6, 7], # 0.05 rad per 20 Hz control step = ~1 rad/s, RATE-matched to r02's # 0.10 rad/step at its 9.2 Hz capture (0.10/step here would be ~2 rad/s # and the wrist visibly whips). step_cap=0.05, servo_caps={6: 0.30, 7: 0.60}, servo_cap_default=0.30, # for listed joints missing from servo_caps j7_total_cap=3.14159, # |dj7_init + servo| (joint 7 only) gripper_mode="threshold", # bit | threshold | continuous | recorded grip_tau=0.35, settle_frames=30, qvel_mode="finitediff", # finitediff | zero embodiment="robocasa", # robocasa (finetuned head) | oxe_droid (BASE model) oxe_grip_flip=False, # oxe_droid: invert the gripper polarity # rel-6D weld (sim_action_aug simfix, user decision 2026-08-06): once the # grasp is REAL (both pads in contact + lift >= weld_lift_m + gripper # commanded closed) the object is welded rigidly to the eef; the weld drops # the moment the gripper is commanded open, so the release is a genuine # fall. Without it a pinned-arm carry slips (old robocasa pin-sweep peaked # at ~33% success; verified again on ep0 2026-08-13). rel6d=True, # engage EARLY (before the object sags in the closing jaw -- a 15 mm gate # let the honey bottle ride 19 mm low and catch the cabinet-shelf edge) and # keep the weld stiff enough that contact cannot stretch it. weld_lift_m=0.004, weld_solref="0.005 1", weld_open_frac=0.5, # "commanded open" = below this fraction of the held cmd weld_open_floor=0.10, # ... or below this absolute close fraction ) WELD_NAME = "actaug_grasp_weld" EEF_BODY = "gripper0_right_eef" def inject_weld(xml, solref="0.005 1"): """Add an INACTIVE weld equality between the eef body and the body that owns the object's free joint (obj_joint0). Idempotent; pure-XML, pre-compile.""" if WELD_NAME in xml: return xml import re # the object root body = the innermost still open at obj_joint0 j = xml.index('name="obj_joint0"') stack = [] for m in re.finditer(r"]*>|", xml[:j]): stack.pop() if m.group(0).startswith("') if "" in xml: return xml.replace("", "" + weld, 1) return xml.replace("", f"{weld}", 1) def _wrap_half_pi(x): """Wrap into (-pi/2, pi/2] -- a 2-jaw gripper is symmetric under 180deg.""" while x > np.pi / 2: x -= np.pi while x <= -np.pi / 2: x += np.pi return x class PinRig: """Joint addresses + pinning helpers for one compiled robocasa model.""" def __init__(self, env): sim = env.env.sim self.sim = sim jn2i = sim.model.joint_name2id self.aq = np.array([int(sim.model.jnt_qposadr[jn2i(f"robot0_joint{i}")]) for i in range(1, 8)]) self.av = np.array([int(sim.model.jnt_dofadr[jn2i(f"robot0_joint{i}")]) for i in range(1, 8)]) self.jids = {i: jn2i(f"robot0_joint{i}") for i in range(1, 8)} # mobile base + torso joints (whatever exists in this model) allj = [sim.model.joint_id2name(j) for j in range(sim.model.njnt)] names = [n for n in allj if n is not None and (n.startswith("mobilebase0_") or "torso" in n)] bq, bv = [], [] for n in names: jid = jn2i(n) bq.append(int(sim.model.jnt_qposadr[jid])) bv.append(int(sim.model.jnt_dofadr[jid])) self.bq = np.array(bq, dtype=int) self.bv = np.array(bv, dtype=int) # gripper finger position-actuators; ctrl=0 = fully closed, open = the # ctrlrange endpoint farthest from 0 (grip_state convention, verified) gca = np.array([int(sim.model.actuator_name2id(a)) for a in ("gripper0_right_gripper_finger_joint1", "gripper0_right_gripper_finger_joint2")]) crange = sim.model.actuator_ctrlrange[gca] self.gca = gca self.grip_close = np.zeros(len(gca)) self.grip_open = np.array( [crange[k][1] if abs(crange[k][1]) >= abs(crange[k][0]) else crange[k][0] for k in range(len(gca))]) ctrl_dt = float(env.env.control_timestep) self.ctrl_dt = ctrl_dt self.n_sub = max(1, int(round(ctrl_dt / sim.model.opt.timestep))) # object geoms + gripper geoms for the contact-gated lift test self.obj_geoms = self._geoms_of_body_prefix("obj") self.grip_geoms = self._geoms_of_body_prefix("gripper0_right") # left/right jaw geoms for the both-pads weld gate self.geoms_L = (self._geoms_of_body_prefix("gripper0_right_leftfinger") | self._geoms_of_body_prefix("gripper0_right_finger_joint1_tip")) self.geoms_R = (self._geoms_of_body_prefix("gripper0_right_rightfinger") | self._geoms_of_body_prefix("gripper0_right_finger_joint2_tip")) # rel-6D weld equality (present only when the XML was inject_weld-ed) import mujoco self._mj = mujoco self._rm = getattr(sim.model, "_model", sim.model) self._rd = getattr(sim.data, "_data", sim.data) self.weld_eq = mujoco.mj_name2id(self._rm, mujoco.mjtObj.mjOBJ_EQUALITY, WELD_NAME) self.b_eef = mujoco.mj_name2id(self._rm, mujoco.mjtObj.mjOBJ_BODY, EEF_BODY) self.b_obj = int(sim.model.jnt_bodyid[jn2i(OBJ_JOINT)]) \ if hasattr(sim.model, "jnt_bodyid") else int(self._rm.jnt_bodyid[ mujoco.mj_name2id(self._rm, mujoco.mjtObj.mjOBJ_JOINT, OBJ_JOINT)]) if self.weld_eq >= 0: self._rd.eq_active[self.weld_eq] = 0 # every attempt starts unwelded @property def weld_active(self): return bool(self.weld_eq >= 0 and self._rd.eq_active[self.weld_eq]) def _quat_conj_mul(self, qa, qb): # qa^-1 * qb, wxyz wa, xa, ya, za = qa qac = np.array([wa, -xa, -ya, -za]) wb, xb, yb, zb = qb return np.array([ qac[0]*wb - qac[1]*xb - qac[2]*yb - qac[3]*zb, qac[0]*xb + qac[1]*wb + qac[2]*zb - qac[3]*yb, qac[0]*yb - qac[1]*zb + qac[2]*wb + qac[3]*xb, qac[0]*zb + qac[1]*yb - qac[2]*xb + qac[3]*wb]) def weld_relpose(self): """Pose of the object body in the eef body frame (p, q wxyz).""" d = self._rd p1, q1 = d.xpos[self.b_eef], d.xquat[self.b_eef] p2, q2 = d.xpos[self.b_obj], d.xquat[self.b_obj] w, x, y, z = q1 R1 = np.array([[1-2*(y*y+z*z), 2*(x*y-z*w), 2*(x*z+y*w)], [2*(x*y+z*w), 1-2*(x*x+z*z), 2*(y*z-x*w)], [2*(x*z-y*w), 2*(y*z+x*w), 1-2*(x*x+y*y)]]) return R1.T @ (p2 - p1), self._quat_conj_mul(q1, q2) def weld_engage(self): """Freeze the CURRENT relative pose and activate the weld (r00_common).""" if self.weld_eq < 0 or self.weld_active: return False p, q = self.weld_relpose() self._rm.eq_data[self.weld_eq, 0:3] = 0.0 self._rm.eq_data[self.weld_eq, 3:6] = p self._rm.eq_data[self.weld_eq, 6:10] = q self._rm.eq_data[self.weld_eq, 10] = 1.0 self._rd.eq_active[self.weld_eq] = 1 self._mj.mj_forward(self._rm, self._rd) return True def weld_release(self): if self.weld_eq < 0 or not self.weld_active: return False self._rd.eq_active[self.weld_eq] = 0 self._mj.mj_forward(self._rm, self._rd) return True def pad_contacts(self): """(left_touching_obj, right_touching_obj).""" d = self.sim.data L = R = False for i in range(d.ncon): c = d.contact[i] g1, g2 = int(c.geom1), int(c.geom2) og = g1 if g1 in self.obj_geoms else (g2 if g2 in self.obj_geoms else None) if og is None: continue other = g2 if og == g1 else g1 if other in self.geoms_L: L = True elif other in self.geoms_R: R = True return L, R def _geoms_of_body_prefix(self, prefix): m = self.sim.model bids = [b for b in range(m.nbody) if (m.body_id2name(b) or "").startswith(prefix)] return set(g for g in range(m.ngeom) if int(m.geom_bodyid[g]) in set(bids)) def grip_contact(self): d = self.sim.data for i in range(d.ncon): c = d.contact[i] g1, g2 = int(c.geom1), int(c.geom2) if (g1 in self.obj_geoms and g2 in self.grip_geoms) or \ (g2 in self.obj_geoms and g1 in self.grip_geoms): return True return False def joint_axis_world(self, i): """World-frame axis of arm joint i (1-indexed) at the current state.""" return np.asarray(self.sim.data.xaxis[self.jids[i]]).ravel().copy() def grip_ctrl(self, frac): """Closure fraction in [0,1] -> per-actuator ctrl (0=open, 1=closed).""" f = float(np.clip(frac, 0.0, 1.0)) return self.grip_open + f * (self.grip_close - self.grip_open) def pin_step(self, q_from, q_to, base_from, base_to, gtarget, qvel_mode): """One control step: ramp the pinned joints across the physics substeps while the gripper actuators + the object integrate freely.""" sim = self.sim arm_vel = (q_to - q_from) / self.ctrl_dt base_vel = (base_to - base_from) / self.ctrl_dt if len(self.bq) else None fd = qvel_mode == "finitediff" for s in range(self.n_sub): f = (s + 1) / self.n_sub sim.data.qpos[self.aq] = q_from + f * (q_to - q_from) sim.data.qvel[self.av] = arm_vel if fd else 0.0 if len(self.bq): sim.data.qpos[self.bq] = base_from + f * (base_to - base_from) sim.data.qvel[self.bv] = base_vel if fd else 0.0 sim.data.ctrl[self.gca] = gtarget sim.step() sim.forward() def _release_frame(actions, grasp_t): """First recorded gripper-open command after the grasp (end of episode if none).""" opens = np.where(actions[grasp_t:, 11] <= 0)[0] return int(grasp_t + opens[0]) if len(opens) else len(actions) - 1 def run_constrained_attempt(env, ep, client, dj7_init_deg, attempt_seed, cfg, record): """One constrained attempt on an ALREADY reset+perturbed env. Returns (row, dump) with the same contract as rollout.run_attempt.""" c = {**DEFAULTS, **(cfg or {})} # `pre_buffer` is the canonical alias for `start_off`: the policy starts at # grasp_t - pre_buffer (same meaning as hybrid mode's --pre_buffer). if "pre_buffer" in c: c["start_off"] = int(c["pre_buffer"]) rig = PinRig(env) sim = rig.sim states = ep["states"] # (T, 1+nq+nv) flattened MjSimState acts = ep["actions"] T = len(states) grasp_t = grasp_frame(acts) rel_t = _release_frame(acts, grasp_t) arm_of = lambda t: states[min(t, T - 1)][1 + rig.aq] base_of = lambda t: states[min(t, T - 1)][1 + rig.bq] if len(rig.bq) \ else np.zeros(0) rec_close = acts[:, 11] > 0 # recorded gripper bit per frame f_pol0 = max(0, grasp_t - int(c["start_off"])) f_hand = min(T - 1, grasp_t + int(c["k_post"])) pj = [int(j) for j in c["policy_joints"]] assert all(1 <= j <= 7 for j in pj), f"policy_joints out of range: {pj}" caps = {int(k): float(v) for k, v in dict(c["servo_caps"]).items()} cap_of = lambda j: caps.get(j, float(c["servo_cap_default"])) # dj7_init (the searched jaw offset) only makes sense when the policy owns # joint 7; otherwise force it to 0 so the recording is untouched. dj7_init = np.radians(float(dj7_init_deg)) if 7 in pj else 0.0 dj7 = 0.0 # ramped in during approach tail djs = {j: 0.0 for j in pj} # per-joint policy servo state init_z = obj_z(env) dump = {"actions": [], "states": [], "frames": [], "obj": []} st = {"succ": False, "peak_lift": 0.0, "lift_grip": 0.0, "policy_steps": 0, "n_infer": 0, "weld_g": None, "n_engage": 0, "engage_t": None, "weld_release_t": None} # executed vs recorded object position per frame (debug track, saved by # save_dump as obj_track.npz when present) jn2i = sim.model.joint_name2id obj_adr = int(sim.model.jnt_qposadr[jn2i(OBJ_JOINT)]) use_weld = bool(c["rel6d"]) and rig.weld_eq >= 0 def rel6d_update(t, g_frac): """Engage the rigid weld on a real grasp; drop it once commanded open.""" if not use_weld: return if rig.weld_active: thr = max(float(c["weld_open_frac"]) * st["weld_g"], float(c["weld_open_floor"])) if g_frac < thr: rig.weld_release() if st["weld_release_t"] is None: st["weld_release_t"] = t return L, R = rig.pad_contacts() dz = obj_z(env) - init_z if L and R and dz >= float(c["weld_lift_m"]) \ and g_frac > float(c["weld_open_floor"]): st["weld_g"] = g_frac rig.weld_engage() st["n_engage"] += 1 if st["engage_t"] is None: st["engage_t"] = t # dj7_init is slewed in over the last frames of the approach (phase R of # r02) instead of teleported, obeying the same per-step cap the servo has. n_ramp = 0 if abs(dj7_init) < 1e-9 else \ int(min(48, max(6, np.ceil(abs(dj7_init) / c["step_cap"])))) n_ramp = min(n_ramp, f_pol0) ramp0 = f_pol0 - n_ramp def track(gbit): st["peak_lift"] = max(st["peak_lift"], obj_z(env) - init_z) if rig.grip_contact(): st["lift_grip"] = max(st["lift_grip"], obj_z(env) - init_z) if is_success(env): st["succ"] = True if record: env.env._get_observations(force_update=True) dump["states"].append(groot_state_vec(env)) pa = env_to_parquet_action(to_env_action(acts[min(len(acts) - 1, len(dump["actions"]))])) pa[11] = 1.0 if gbit else -1.0 dump["actions"].append(pa) dump["frames"].append(render3(env)) t_rec = min(len(dump["actions"]) - 1, T - 1) dump["obj"].append(np.concatenate([ sim.data.qpos[obj_adr:obj_adr + 3].copy(), states[t_rec][1 + obj_adr:1 + obj_adr + 3]])) q_prev = arm_of(0).copy() b_prev = base_of(0).copy() sim.data.qpos[rig.aq] = q_prev if len(rig.bq): sim.data.qpos[rig.bq] = b_prev sim.forward() def frame_to(t, dj, gfrac, gbit): """dj = {joint_index_1based: offset_rad} applied on top of the recording.""" nonlocal q_prev, b_prev q_to = arm_of(t).copy() for j, v in dj.items(): q_to[j - 1] += v b_to = base_of(t).copy() rig.pin_step(q_prev, q_to, b_prev, b_to, rig.grip_ctrl(gfrac), c["qvel_mode"]) q_prev, b_prev = q_to, b_to rel6d_update(t, gfrac) track(gbit) # ---- phase 0: pinned approach (recorded gripper bit), with the dj7 ramp for t in range(0, f_pol0): if t >= ramp0 and n_ramp: dj7 = dj7_init * (t - ramp0 + 1) / n_ramp g = bool(rec_close[min(t, len(acts) - 1)]) frame_to(t, {7: dj7} if 7 in pj else {}, 1.0 if g else 0.0, g) dj7 = dj7_init if 7 in pj: djs[7] = dj7 # ---- phase A: policy joint servo (policy_joints) + gripper chunk, ci = None, 0 grip_latch = False g_frac = 1.0 if rec_close[max(0, f_pol0 - 1)] else 0.0 grip_hist = [g_frac] no_policy = c["gripper_mode"] == "recorded" and \ all(cap_of(j) == 0.0 for j in pj) for t in range(f_pol0, f_hand + 1): if st["succ"]: break if no_policy: # pure grip_state diagnostic: no inference at all g = bool(rec_close[min(t, len(acts) - 1)]) g_frac = 1.0 if g else 0.0 grip_latch = grip_latch or g grip_hist.append(g_frac) frame_to(t, dict(djs), g_frac, g) continue oxe = c["embodiment"] == "oxe_droid" if chunk is None or ci >= int(c["replan"]): env.env._get_observations(force_update=True) seed = attempt_seed * 100003 + st["n_infer"] obs_fn = groot_obs_oxe if oxe else groot_obs chunk = client.get_action(obs_fn(env, ep["instruction"], action_seed=seed)) st["n_infer"] += 1 ci = 0 if oxe: a_rot = np.asarray(chunk["action.eef_rotation_delta"])[ci].ravel()[:3] g = float(np.asarray(chunk["action.gripper_position"])[ci].ravel()[0]) a_g = 1.0 - g if c["oxe_grip_flip"] else g # close fraction in [0,1] else: a_rot = np.asarray(chunk["action.end_effector_rotation"])[ci].ravel()[:3] a_g = float(np.asarray(chunk["action.gripper_close"])[ci].ravel()[0]) ci += 1 st["policy_steps"] += 1 # delta-rotation reduced to its component about the wrist axis. # The delta is base-frame axis-angle; the base yaw is pinned to the # recording, so rotate it into the world frame before projecting. di = env.env._get_observations(force_update=False) bq = np.asarray(di["robot0_base_quat"]).ravel() # xyzw x, y, z, w = bq Rb = np.array([[1-2*(y*y+z*z), 2*(x*y-z*w), 2*(x*z+y*w)], [2*(x*y+z*w), 1-2*(x*x+z*z), 2*(y*z-x*w)], [2*(x*z-y*w), 2*(y*z+x*w), 1-2*(x*x+y*y)]]) rv_w = Rb @ a_rot sc = float(c["step_cap"]) for j in pj: e = float(np.dot(rv_w, rig.joint_axis_world(j))) if j == 7: # jaw 180deg symmetry is about the roll axis only e = _wrap_half_pi(e) d = float(np.clip(e, -sc, sc)) servo = float(np.clip(dj7 + d - dj7_init, -cap_of(7), cap_of(7))) dj7 = float(np.clip(dj7_init + servo, -c["j7_total_cap"], c["j7_total_cap"])) djs[7] = dj7 else: d = float(np.clip(e, -sc, sc)) djs[j] = float(np.clip(djs[j] + d, -cap_of(j), cap_of(j))) mode = c["gripper_mode"] if mode == "threshold": if a_g >= float(c["grip_tau"]): grip_latch = True g_frac = 1.0 if grip_latch else 0.0 elif mode == "continuous": g_frac = float(np.clip(a_g, 0.0, 1.0)) elif mode == "recorded": # diagnostic: the recorded schedule, not the policy g_frac = 1.0 if rec_close[min(t, len(acts) - 1)] else 0.0 grip_latch = grip_latch or g_frac >= 1.0 else: # bit: the binarized robocasa-native command g_frac = 1.0 if a_g > 0 else 0.0 grip_hist.append(a_g) frame_to(t, dict(djs), g_frac, g_frac >= 0.5) # ---- hand-off: freeze every joint offset + hold the achieved close command dj_hold = dict(djs) g_hold = (1.0 if grip_latch else float(np.clip(max(grip_hist[-3:]), 0.0, 1.0))) \ if c["gripper_mode"] != "bit" else g_frac # ---- phase B: pinned transport, recorded opening schedule from rel_t for t in range(f_hand + 1, T): if st["succ"]: break if t >= rel_t: g = min(g_hold, 1.0 if rec_close[min(t, len(acts) - 1)] else 0.0) else: g = g_hold frame_to(t, dj_hold, g, g >= 0.5) # ---- settle: hold the last pose so the final pose is a rest pose g_end = 0.0 if T - 1 >= rel_t else g_hold for _ in range(int(c["settle_frames"])): if st["succ"]: break rig.pin_step(q_prev, q_prev, b_prev, b_prev, rig.grip_ctrl(g_end), c["qvel_mode"]) rel6d_update(T - 1, g_end) track(False) row = dict(mode="constrained", attempt_seed=attempt_seed, grasp_t=grasp_t, release_t=rel_t, dj7_init_deg=float(dj7_init_deg), policy_joints=str(pj), dj_handoff_deg=" ".join(f"{j}:{np.degrees(v):+.1f}" for j, v in sorted(dj_hold.items())), grip_handoff=round(float(g_hold), 3), gripper_mode=c["gripper_mode"], n_infer=st["n_infer"], policy_steps=st["policy_steps"], peak_lift_m=round(st["peak_lift"], 4), lift_grip_m=round(st["lift_grip"], 4), rel6d=int(use_weld), weld_engage_t=st["engage_t"], weld_release_t=st["weld_release_t"], n_weld_engage=st["n_engage"], grasped=int(st["lift_grip"] > LIFT_M), success=int(st["succ"]), n_steps=len(dump["actions"]) if record else -1) return row, dump