Files
Mujoco_WASM/training_server/tests/test_mobile_manipulator.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

225 lines
8.7 KiB
Python

import copy
import importlib.util
import json
import tempfile
import unittest
from pathlib import Path
from xml.etree import ElementTree as ET
import numpy as np
from training_server.mobile_manipulator.kernel import (
CONTRACTS,
ROBOTS,
TASK,
TaskKernel,
decode_action,
)
def make_package(path):
import mujoco
config = copy.deepcopy(ROBOTS[0])
config["recipe"] = "mjcf"
root = ET.Element("mujoco")
ET.SubElement(root, "compiler", angle="radian")
ET.SubElement(root, "option", timestep=".002", gravity="0 0 0", integrator="implicitfast")
world = ET.SubElement(root, "worldbody")
base = ET.SubElement(world, "body", name=config["baseBodyName"], pos="0 0 .1")
ET.SubElement(base, "freejoint", name=config["baseJointName"])
ET.SubElement(base, "geom", type="sphere", size=".05", mass="2", contype="0", conaffinity="0")
actuators = ET.SubElement(root, "actuator")
for name in config["baseJoints"]:
body = ET.SubElement(base, "body")
ET.SubElement(body, "joint", name=name, axis="0 1 0", damping=".1")
ET.SubElement(
body, "geom", type="sphere", size=".02", mass=".1", contype="0", conaffinity="0"
)
ET.SubElement(
actuators,
"velocity",
name=name + "_servo",
joint=name,
kv="1",
ctrllimited="true",
ctrlrange=f"{-config['wheelLimit']} {config['wheelLimit']}",
)
for j in [
*config["armJoints"],
dict(name=config["gripperJoint"], min=config["gripperClosed"], max=config["gripperOpen"]),
]:
body = ET.SubElement(base, "body", pos="0 0 .1")
if j["name"] == config["armJoints"][-1]["name"]:
body.set("name", config["eefBodyName"])
ET.SubElement(body, "site", name=config["eefSiteName"], pos=".1 0 0")
ET.SubElement(
body,
"joint",
name=j["name"],
axis="0 0 1",
limited="true",
range=f"{j['min']} {j['max']}",
damping=".1",
armature=".01",
)
ET.SubElement(
body, "geom", type="sphere", size=".02", mass=".1", contype="0", conaffinity="0"
)
ET.SubElement(
actuators,
"position",
name=j["name"] + "_servo",
joint=j["name"],
kp="10",
kv="1",
ctrllimited="true",
ctrlrange=f"{j['min']} {j['max']}",
)
body = ET.SubElement(
world, "body", name="__mm_object", pos=" ".join(map(str, TASK["objectStart"]))
)
ET.SubElement(body, "freejoint", name="__mm_object_joint")
ET.SubElement(body, "geom", type="box", size=".018 .018 .018", mass=".04")
ET.SubElement(
world, "body", name="__mm_goal", mocap="true", pos=" ".join(map(str, TASK["goalStart"]))
)
(path / "scene.xml").write_text(ET.tostring(root, encoding="unicode"))
(path / "robot.json").write_text(json.dumps(config, separators=(",", ":")))
(path / "task.json").write_text(json.dumps(TASK))
(path / "environment.json").write_text(
json.dumps(
dict(
scene="scene.xml",
mujoco=mujoco.__version__,
taskId=TASK["id"],
robotId=config["id"],
)
)
)
class KernelTests(unittest.TestCase):
def test_golden_and_reusable_buffers(self):
cases = json.loads((CONTRACTS / "fixtures/mobile-golden.json").read_text())
for case in cases:
config = next(c for c in ROBOTS if c["id"] == case["robotId"])
kernel = TaskKernel(config)
kernel.has_lifted = case["lifted"]
observation, reward, *_, info = kernel.evaluate(np.array(case["state"]))
np.testing.assert_allclose(observation, case["observation"], atol=1e-7)
np.testing.assert_allclose(
decode_action(config, case["action"]), case["control"], atol=1e-12
)
self.assertAlmostEqual(reward, case["reward"], places=12)
self.assertEqual(info["stage"], case["info"]["stage"])
self.assertIs(kernel.observe(np.array(case["state"])), observation)
def test_invalid_action_is_atomic_and_velocity_mode(self):
config = copy.deepcopy(ROBOTS[0])
config["armJoints"][0]["mode"] = "velocity"
output = np.ones(9)
action = np.zeros(12, dtype=np.float32)
action[3] = 0.5
decode_action(config, action, output)
self.assertEqual(output[3], 1)
previous = output.copy()
action[11] = np.nan
with self.assertRaises(ValueError):
decode_action(config, action, output)
np.testing.assert_equal(previous, output)
@unittest.skipUnless(
importlib.util.find_spec("mujoco") and importlib.util.find_spec("gymnasium"),
"install mobile_manipulator/requirements.txt for native tests",
)
class NativeTests(unittest.TestCase):
def setUp(self):
from training_server.mobile_manipulator.env import MobileManipulatorEnv
self.directory = tempfile.TemporaryDirectory()
self.path = Path(self.directory.name)
make_package(self.path)
self.env = MobileManipulatorEnv(self.path)
def tearDown(self):
self.env.close()
self.directory.cleanup()
def test_gym_contract_and_reset(self):
from gymnasium.utils.env_checker import check_env
check_env(self.env, skip_render_check=True)
obs, _ = self.env.reset(seed=42)
action = self.env.action.copy()
action[3] = 0.2
stepped = self.env.step(action)
self.assertTrue(np.isfinite(stepped[0]).all())
self.assertFalse(np.shares_memory(obs, stepped[0]))
self.assertAlmostEqual(self.env.data.time, TASK["controlDt"])
self.env.move_task_entity("goal", [0.7, 0.2, 0.019])
np.testing.assert_allclose(self.env.state[47:50], [0.7, 0.2, 0.019])
self.assertEqual(self.env.kernel.steps, 0)
self.env.kernel.has_lifted = True
self.env.move_task_entity("object", [0.3, 0.1, 0.02])
self.assertFalse(self.env.kernel.has_lifted)
self.assertEqual(self.env.data.qvel[self.env.object_v : self.env.object_v + 6].sum(), 0)
self.env.close()
with self.assertRaises(RuntimeError):
self.env.step(action)
def test_navigation_holds_arm_and_has_observable_controller_state(self):
self.env.kernel.stage = "navigate"
self.env.reset(seed=7)
before = self.env.control.copy()
obs, _, _, _, info = self.env.step(np.ones(TASK["actionSize"], dtype=np.float32))
np.testing.assert_allclose(self.env.control[3:], before[3:])
np.testing.assert_equal(obs[68:80], self.env.motion.applied)
np.testing.assert_equal(obs[80:92], self.env.motion.targets)
self.assertEqual(info["safety_stop"], "")
self.assertLess(info["max_joint_velocity"], TASK["jointSpeedStop"])
self.env.data.qvel[self.env.arm[0][1]] = 30
self.assertTrue(self.env.step(np.zeros(12, dtype=np.float32))[2])
self.assertEqual(self.env.kernel.info["safety_stop"], "joint_velocity")
def test_hold_does_not_ratchet_targets_toward_gravity_sag(self):
import mujoco
previous = self.env.control.copy()
self.env.data.qpos[self.env.arm[0][0]] -= 0.1
mujoco.mj_forward(self.env.model, self.env.data)
for _ in range(10):
self.env.hold()
np.testing.assert_equal(self.env.control, previous)
def test_randomized_reset_is_seeded_and_evaluation_can_be_fixed(self):
self.env.position_jitter = 0.1
first, _ = self.env.reset(seed=7)
second, _ = self.env.reset(seed=7)
third, _ = self.env.reset(seed=8)
np.testing.assert_equal(first, second)
self.assertFalse(np.array_equal(first, third))
self.env.position_jitter = 0
first, _ = self.env.reset(seed=7)
second, _ = self.env.reset(seed=8)
np.testing.assert_equal(first, second)
def test_reject_mismatching_model_and_version(self):
from training_server.mobile_manipulator.env import MobileManipulatorEnv
meta = json.loads((self.path / "environment.json").read_text())
meta["mujoco"] = "0.0.0"
(self.path / "environment.json").write_text(json.dumps(meta))
with self.assertRaisesRegex(ValueError, "version mismatch"):
MobileManipulatorEnv(self.path)
config = json.loads((self.path / "robot.json").read_text())
config["armJoints"][0]["max"] = 1
(self.path / "robot.json").write_text(json.dumps(config))
with self.assertRaisesRegex(ValueError, "joint range"):
MobileManipulatorEnv(self.path, allow_version_mismatch=True)
if __name__ == "__main__":
unittest.main()