"""Native Gymnasium side. Load the *browser-exported* scene, not the source URDF. python -m training_server.mobile_manipulator.export_onnx --help """ import copy import json from pathlib import Path import gymnasium as gym import mujoco import numpy as np from .kernel import TASK, TaskKernel, clip, rotate, validate_config from .motion import SafeActionController class MobileManipulatorEnv(gym.Env): metadata = {"render_modes": []} def __init__( self, package, *, allow_version_mismatch=False, reset_options=None, stage="pick-place", position_jitter=0.0, ): super().__init__() self.reset_options = copy.deepcopy(reset_options) self.position_jitter = position_jitter directory = Path(package).resolve() metadata = json.loads((directory / "environment.json").read_text()) scene = (directory / metadata["scene"]).resolve() if not scene.is_relative_to(directory): raise ValueError("scene escapes package") if metadata["mujoco"] != mujoco.__version__ and not allow_version_mismatch: raise ValueError( f"MuJoCo version mismatch: package={metadata['mujoco']}, " f"native={mujoco.__version__}" ) if json.loads((directory / "task.json").read_text()) != TASK: raise ValueError("task contract mismatch") self.config = json.loads((directory / "robot.json").read_text()) validate_config(self.config) if metadata["taskId"] != TASK["id"] or metadata["robotId"] != self.config["id"]: raise ValueError("package identity mismatch") self.model = mujoco.MjModel.from_xml_path(str(scene)) self.data = mujoco.MjData(self.model) self.kernel = TaskKernel(self.config, stage) self.motion = SafeActionController(self.config) self.state = np.zeros(TASK["stateSize"], dtype=np.float64) self.action_space = gym.spaces.Box(-1, 1, (TASK["actionSize"],), dtype=np.float32) self.observation_space = gym.spaces.Box(-1, 1, (TASK["observationSize"],), dtype=np.float32) self.frame_skip = round(TASK["controlDt"] / self.model.opt.timestep) if ( self.frame_skip < 1 or abs(self.frame_skip * self.model.opt.timestep - TASK["controlDt"]) > 1e-9 ): raise ValueError("controlDt must be a multiple of timestep") self.base_q, self.base_v = self._free(self.config["baseJointName"]) base_joint = self._id(mujoco.mjtObj.mjOBJ_JOINT, self.config["baseJointName"]) if self.model.jnt_bodyid[base_joint] != self._id( mujoco.mjtObj.mjOBJ_BODY, self.config["baseBodyName"] ): raise ValueError("base joint/body mismatch") self.object_q, self.object_v = self._free("__mm_object_joint") self.eef_body = self._id(mujoco.mjtObj.mjOBJ_BODY, self.config["eefBodyName"]) self.eef_site = ( self._id(mujoco.mjtObj.mjOBJ_SITE, self.config["eefSiteName"]) if self.config.get("eefSiteName") else -1 ) self.goal_mocap = self.model.body_mocapid[self._id(mujoco.mjtObj.mjOBJ_BODY, "__mm_goal")] if self.goal_mocap < 0: raise ValueError("goal must be mocap") self.arm = [self._scalar(j["name"]) for j in self.config["armJoints"]] for spec in self.config["armJoints"]: j = self._id(mujoco.mjtObj.mjOBJ_JOINT, spec["name"]) if not self.model.jnt_limited[j] or not np.allclose( self.model.jnt_range[j], [spec["min"], spec["max"]], atol=1e-5, rtol=0 ): raise ValueError(f"joint range mismatch: {spec['name']}") self.grippers = [ self._scalar(g.get("joint", self.config["gripperJoint"])) for g in self.config["gripperActuators"] ] self.gripper_q = self._scalar(self.config["gripperJoint"])[0] bindings = [] for name, joint in zip( self.config["baseActuators"], self.config["baseJoints"], strict=True ): bindings.append( self._actuator( name, joint, "velocity", -self.config["wheelLimit"], self.config["wheelLimit"] ) ) for name, j in zip(self.config["armActuators"], self.config["armJoints"], strict=True): lo, hi = ( (j["min"], j["max"]) if j["mode"] == "position" else (-j["velocityLimit"], j["velocityLimit"]) ) bindings.append(self._actuator(name, j["name"], j["mode"], lo, hi)) for g in self.config["gripperActuators"]: bindings.append( self._actuator( g["name"], g.get("joint", self.config["gripperJoint"]), "position", min(g["closed"], g["open"]), max(g["closed"], g["open"]), ) ) self.control_addresses = np.array(bindings, dtype=int) self.control = self.motion.control self.action = np.zeros(TASK["actionSize"], dtype=np.float32) self._closed = False self.reset() def _id(self, kind, name): result = mujoco.mj_name2id(self.model, kind, name) if result < 0: raise ValueError(f"missing model name: {name}") return result def _free(self, name): i = self._id(mujoco.mjtObj.mjOBJ_JOINT, name) if self.model.jnt_type[i] != mujoco.mjtJoint.mjJNT_FREE: raise ValueError(f"{name} must be freejoint") return self.model.jnt_qposadr[i], self.model.jnt_dofadr[i] def _scalar(self, name): i = self._id(mujoco.mjtObj.mjOBJ_JOINT, name) if self.model.jnt_type[i] not in (mujoco.mjtJoint.mjJNT_HINGE, mujoco.mjtJoint.mjJNT_SLIDE): raise ValueError(f"{name} must be scalar joint") return self.model.jnt_qposadr[i], self.model.jnt_dofadr[i] def _actuator(self, name, joint, mode, lo, hi): m = self.model i = self._id(mujoco.mjtObj.mjOBJ_ACTUATOR, name) j = self._id(mujoco.mjtObj.mjOBJ_JOINT, joint) self._scalar(joint) addresses = getattr(m, "actuator_ctrladr", np.arange(m.nu)) address = addresses[i] end = addresses[i + 1] if i + 1 < len(addresses) else m.nu gain, bp = m.actuator_gainprm[i, 0], m.actuator_biasprm[i] if ( end - address != 1 or m.actuator_trntype[i] != 0 or m.actuator_trnid[i, 0] != j or m.actuator_gaintype[i] != 0 or m.actuator_dyntype[i] != 0 or gain <= 0 or abs(m.actuator_gear[i, 0] - 1) > 1e-8 or m.actuator_biastype[i] != 1 or ( abs(bp[1] + gain) > 1e-6 if mode == "position" else abs(bp[1]) > 1e-8 or abs(bp[2] + gain) > 1e-6 ) or not m.actuator_ctrllimited[i] or not np.allclose(m.actuator_ctrlrange[i], [lo, hi], atol=1e-5, rtol=0) ): raise ValueError(f"actuator contract mismatch: {name}") return address def _check(self): if self._closed: raise RuntimeError("environment closed") def reset(self, *, seed=None, options=None): self._check() super().reset(seed=seed) mujoco.mj_resetData(self.model, self.data) self.action.fill(0) for (q, _), j in zip(self.arm, self.config["armJoints"], strict=True): self.data.qpos[q] = j["neutral"] self.data.qpos[self.gripper_q] = self.config["gripperOpen"] for (q, _), g in zip(self.grippers, self.config["gripperActuators"], strict=True): self.data.qpos[q] = g["open"] self.kernel.reset() mujoco.mj_forward(self.model, self.data) self.hold(preserve_targets=False) # Training-only seeded sampling; evaluation uses its own fixed seed sequence. # Browser playback uses the deployment's nominal reset, never a hidden PRNG. options = copy.deepcopy(self.reset_options if options is None else options) if self.position_jitter: options = options or { "object": TASK["objectStart"].copy(), "goal": TASK["goalStart"].copy(), } for entity in ("object", "goal"): if entity in options: for i in range(2): options[entity][i] = float( np.clip( options[entity][i] + self.np_random.uniform( -self.position_jitter, self.position_jitter ), -TASK["positionScale"] + 0.1, TASK["positionScale"] - 0.1, ) ) if options: for entity in ("object", "goal"): if entity in options: self.move_task_entity(entity, options[entity]) return self.observe().copy(), copy.deepcopy(self.kernel.info) def _apply(self, action): self.motion.apply(action, self.state, self.kernel.stage, self.kernel.has_lifted) self.kernel.record_action(self.motion.applied, self.motion.targets) self.data.ctrl[self.control_addresses] = self.control def observe(self): self._check() s, d = self.state, self.data s[:7] = d.qpos[self.base_q : self.base_q + 7] s[7:10] = d.qvel[self.base_v : self.base_v + 3] s[10:13] = rotate(s[3:7], d.qvel[self.base_v + 3 : self.base_v + 6]) for i, (q, v) in enumerate(self.arm): s[13 + i], s[21 + i] = d.qpos[q], d.qvel[v] s[29] = clip( (d.qpos[self.gripper_q] - self.config["gripperClosed"]) / (self.config["gripperOpen"] - self.config["gripperClosed"]), 0, 1, ) if self.eef_site >= 0: s[30:33] = d.site_xpos[self.eef_site] mujoco.mju_mat2Quat(s[33:37], d.site_xmat[self.eef_site]) else: s[30:33], s[33:37] = d.xpos[self.eef_body], d.xquat[self.eef_body] s[37:44] = d.qpos[self.object_q : self.object_q + 7] s[44:47] = d.qvel[self.object_v : self.object_v + 3] s[47:50], s[50:54] = d.mocap_pos[self.goal_mocap], d.mocap_quat[self.goal_mocap] if not np.isfinite(s).all(): raise RuntimeError("non-finite simulation state") return self.kernel.observe(s) def step(self, action): self._check() if self.kernel.terminated or self.kernel.truncated: raise RuntimeError("episode ended; reset required") self._apply(action) peak = 0.0 safety = "" for _ in range(self.frame_skip): mujoco.mj_step(self.model, self.data) peak = max(peak, max(abs(self.data.qvel[v]) for _, v in [*self.arm, *self.grippers])) if peak > TASK["jointSpeedStop"]: safety = "joint_velocity" break mujoco.mj_forward(self.model, self.data) self.observe() obs, reward, terminated, truncated, info = self.kernel.evaluate(self.state, safety, peak) return obs.copy(), reward, terminated, truncated, copy.deepcopy(info) def hold(self, preserve_targets=True): self._check() self.action.fill(0) self.observe() self.motion.reset( self.state, self.data.ctrl[self.control_addresses] if preserve_targets else None ) self.kernel.record_action(self.motion.applied, self.motion.targets) self.data.ctrl[self.control_addresses] = self.control def move_task_entity(self, entity, position): self._check() p = np.asarray(position, dtype=float) if p.shape != (3,) or not np.isfinite(p).all() or max(abs(p)) > TASK["positionScale"]: raise ValueError("invalid task position") p = p.copy() p[2] = max(TASK["objectStart"][2], p[2]) if entity == "object": self.data.qpos[self.object_q : self.object_q + 3] = p self.data.qvel[self.object_v : self.object_v + 6] = 0 elif entity == "goal": self.data.mocap_pos[self.goal_mocap] = p else: raise ValueError("entity must be object or goal") self.data.qacc_warmstart.fill(0) self.kernel.reset() mujoco.mj_forward(self.model, self.data) self.hold() self.observe() def close(self): if not self._closed: self._closed = True self.data = self.model = None # Python bindings own native objects; release references.