Files
Mujoco_WASM/training_server/mobile_manipulator/env.py
T
chenlin f3a8a38acd
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
2026-09-24 09:57:41 +08:00

297 lines
13 KiB
Python

"""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.