actaug: simulator action-augmentation generation code, Table-2 eval protocol, raw dumps of the 256 prio episodes
bd08baf verified Download actaug/code/constrained.py from Ronaldo-GOAT/transfer: direct link, hf CLI and curl.
- Browser
- Download file 24.6 kB
-
https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/actaug/code/constrained.py
- Command line
-
hf download hf://Ronaldo-GOAT/transfer/actaug/code/constrained.py
-
curl -L -o constrained.py https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/actaug/code/constrained.py
24.6 kB
| #!/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 <body> still open at obj_joint0 | |
| j = xml.index('name="obj_joint0"') | |
| stack = [] | |
| for m in re.finditer(r"<body\b[^>]*>|</body>", xml[:j]): | |
| stack.pop() if m.group(0).startswith("</") else stack.append(m.group(0)) | |
| obj_body = re.search(r'name="([^"]+)"', stack[-1]).group(1) | |
| weld = (f'<weld name="{WELD_NAME}" body1="{EEF_BODY}" body2="{obj_body}" ' | |
| f'active="false" solref="{solref}"/>') | |
| if "<equality>" in xml: | |
| return xml.replace("<equality>", "<equality>" + weld, 1) | |
| return xml.replace("</mujoco>", f"<equality>{weld}</equality></mujoco>", 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 | |
| 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 | |