feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
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

集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
This commit is contained in:
2026-09-24 09:57:41 +08:00
parent 3ad29356c9
commit f3a8a38acd
194 changed files with 32918 additions and 236 deletions
+30
View File
@@ -4,6 +4,34 @@
仓库已在 [`rl/`](rl/) 内置 `Unitree-Go2-Flat` 所需的 PPO 训练代码、Go2 模型资产和 ONNX 导出逻辑,不再要求另外克隆 `unitree_rl_mjlab`。`mjlab`、PyTorch 等大型运行依赖仍需安装在本机训练环境中。
## LeKiwi 一键训练(与 Go2 共用控制台)
注册任务 `MobileManipulator-LeKiwi-v1`、`MobileManipulator-LeKiwi-Bundle`,使用 SB3 PPO + 原生 MuJoCo,沿用同一作业状态机、资源锁、取消和日志轮询。MuJoCo 必须与浏览器一致(3.11.0),**不要升级现有 Go2 环境**;首次配置独立解释器:
```bash
source .venv/bin/activate
python -m venv build/venvs/mobile
build/venvs/mobile/bin/python -m pip install -r training_server/mobile_manipulator/requirements.txt
python training_server/server.py --mobile-python "$PWD/build/venvs/mobile/bin/python"
```
主工作台导入对应 ZIP,在 URDF 选项选择「LeKiwi v1 / Bundle · 移动操作训练场景」;旧 `lekiwi-v1` profile 也能从训练面板开始,届时自动组合任务场景。进入「控制台 → 强化学习任务」,连接服务、配置迭代/环境数/设备/种子/每环境采样步数/物体和目标坐标,点击「发起本地训练」。无需下载训练包或手动执行训练脚本。加载场景会暂停仿真并切换到移动操作任务(不启用外部控制桥、Go2 地图或相机配置)。
- 浏览器自动上传组合后的 MJCF 与资产快照。认证 ZIP 上传限 128 MiB、展开 512 MiB/10000 文件;拒绝路径越界、符号链接、include/plugin、非注册机器人契约和包外资产引用。最多保留 20 份去重快照,满额后需停止服务再清理 `logs/mobile_packages/`。
- `POST /jobs` 仅接收服务器内容 ID `mobilePackageId`,不接受本机路径。`mobileParams` 包括 `stage`(默认 `navigate`)、`sourceJobId`、`navigationBootstrapSteps`(0–10000,默认4096)、`positionJitter`(0–0.3 m,默认0.1)、`evaluationEpisodes`(2–64,默认10)、`rolloutSteps`(8–4096,默认128)、`objectPosition`、`goalPosition`。环境数1–64,默认1;总采样步数为 `maxIterations × numEnvs × rolloutSteps`。CPU 物理向量环境顺序采样,设备选择只控制 PPO 网络;GPU 仅支持单卡,不会静默回退 CPU。
- actor/critic/entropy 损失、每采样步平均 reward 和已完成回合长度使用共享 `Learning iteration / Mean ...` 日志协议。没有额外 VecNormalize。
- 自动生成 `logs/rsl_rl/web_jobs/{id}/policy.onnx`、`deployment.json`、PPO checkpoint 与训练配置。元数据包含任务/变体、固定 float32 `[1,92] → [1,12]`、v2 限速动作语义、训练阶段、独立评估、50Hz、权重/机器人配置/场景 SHA-256 和初始目标坐标。
- 点击「导入策略」自动下载两份成果物,经校验后通过主仿真会话的 `ONNXPolicyRunner` 执行单飞锁步推断。暂停、重载和控制权切换使在途结果失效;回合终止/超时暂停,需重置后继续。改动资产/动力学后应重新训练,哈希不是跨机器人泛化证明。
- `health.taskMetadata` 按任务报告 `ready/error`,移动依赖缺失不阻断 Go2,反之亦然。短程冒烟只验证链路,不保证抓取成功率;尚未验证长程收敛、GPU 性能或实机迁移。
阶段按底盘接近→末端接近→抓取放置推进;服务验证上一阶段至少10回合、≥80%成功率且无安全终止。同阶段可继续训练;接续只能引用同一服务会话中已完成的同场景作业。完整设计、限速说明和实际导航结果见 [分阶段训练](../docs/mobile-training-curriculum.md)。旧68维策略必须重新训练。
真实浏览器端到端(需两个本地模型 ZIP 和上述独立环境):
```bash
npx playwright test -c web_platform/playwright.lekiwi.config.ts lekiwi.training.spec.ts
```
## 自定义任务与地形
内置新增 `Unitree-Go2-ObstacleAvoidance`(81维前视射线导航),并放行 `Unitree-Go2-Rough` 训练。健康接口提供可配置参数元数据,job 的 `deployment` 返回精确地图布局、传感器及策略契约。首版支持 `plane/discrete_obstacles/rough/pyramid_stairs/wave` 的训练专用box布局;不是任意场景导入,也不实现真实深度相机。旧 Rough 的234维actor不能在当前浏览器一键部署。完整字段、坐标系、观测动作及复现方式见 [避障部署契约](OBSTACLE_AVOIDANCE.md)。
@@ -76,10 +104,12 @@ python training_server/server.py \
- `GET /api/training/health`:运行环境、允许的任务和活动任务;
- `POST /api/training/pretrained-sources/upload?format=pt|onnx&template=go2-legacy47-v1&name=显示名`:认证有界二进制单文件上传,返回持久化内容ID;
- `POST /api/training/mobile-packages`:认证场景快照上传(`application/zip`),返回64位内容ID;
- `POST /api/training/jobs`:发起训练;
- `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志;
- `DELETE /api/training/jobs/{id}`:停止训练;
- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略;
- `GET /api/training/jobs/{id}/artifacts/deployment.json`:下载移动操作部署元数据;
- `GET /api/tuning/capabilities`、`POST /api/tuning/agent/test`:检查/测试 Agent;
- `GET|POST /api/tuning/sessions`、`GET|DELETE /api/tuning/sessions/{id}`:列出、创建、查询、停止 session;
- `POST /api/tuning/sessions/{id}/pause|resume`:暂停后续调度或恢复;
@@ -0,0 +1 @@
"""Mobile manipulation v1: shared contract, native MuJoCo environment and ONNX export."""
@@ -0,0 +1,84 @@
"""Optional navigation-only behavior-cloning initialization, followed by PPO.
The teacher is used ONLY to collect training data. Export contains the trained MLP,
not a hidden scripted navigation fallback. Later stages never use this initializer.
"""
import math
import mujoco
import numpy as np
import torch
from .env import MobileManipulatorEnv
from .kernel import TASK, navigation_error
def navigation_teacher(state):
_, yaw = navigation_error(state)
x = state[37] - TASK["navigationOffset"][0] - state[0]
y = state[38] - TASK["navigationOffset"][1] - state[1]
action = np.zeros(TASK["actionSize"], dtype=np.float32)
action[0] = 8 * (math.cos(yaw) * x + math.sin(yaw) * y)
action[1] = 8 * (-math.sin(yaw) * x + math.cos(yaw) * y)
action[2] = -4 * yaw
return np.clip(action, -1, 1)
def bootstrap_navigation(agent, package, params, reset_options, seed):
steps = params["navigationBootstrapSteps"]
if not steps:
return
print(f"Navigation BC initialization: {steps} simulated teacher steps, then PPO", flush=True)
env = MobileManipulatorEnv(
package,
stage="navigate",
reset_options=reset_options,
position_jitter=params["positionJitter"],
)
rng = np.random.default_rng(seed)
observations, actions = [], []
def reset():
env.reset()
# Cover heading errors rather than cloning only a straight, zero-yaw path.
yaw = rng.uniform(-0.6, 0.6)
env.data.qpos[env.base_q + 3 : env.base_q + 7] = [
math.cos(yaw / 2),
0,
0,
math.sin(yaw / 2),
]
mujoco.mj_forward(env.model, env.data)
env.hold()
return env.observe().copy()
try:
env.reset(seed=seed)
observation = reset()
for _ in range(steps):
action = navigation_teacher(env.state)
observations.append(observation.copy())
actions.append(action)
# Cover nearby off-teacher states without teleportation or privileged deployment inputs.
noisy = action.copy()
noisy[:3] += rng.normal(0, 0.2, 3)
observation, _, terminated, truncated, _ = env.step(noisy)
if terminated or truncated:
observation = reset()
finally:
env.close()
x = torch.as_tensor(np.asarray(observations), device=agent.device)
y = torch.as_tensor(np.asarray(actions), device=agent.device)
policy = agent.policy
optimizer = torch.optim.Adam(policy.parameters(), lr=1e-3)
for _ in range(30):
for batch in torch.randperm(steps, device=agent.device).split(256):
features = policy.extract_features(x[batch])
predicted = policy.action_net(policy.mlp_extractor.forward_actor(features))
loss = torch.nn.functional.mse_loss(predicted, y[batch])
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(policy.parameters(), 1)
optimizer.step()
print(f"Navigation BC initialization finished; final batch MSE={loss.item():.6g}", flush=True)
@@ -0,0 +1,93 @@
"""Dependency-free task registry and bounded PPO settings."""
import json
import math
import re
from pathlib import Path
CONTRACTS = Path(__file__).resolve().parents[2] / "contracts"
TASK = json.loads((CONTRACTS / "mobile-manipulator-v2.json").read_text())
ROBOTS = {r["id"]: r for r in json.loads((CONTRACTS / "mobile-robots-v1.json").read_text())}
MOBILE_TASKS = {
"MobileManipulator-LeKiwi-v1": "lekiwi-v1",
"MobileManipulator-LeKiwi-Bundle": "lekiwi-bundle",
}
def validate_mobile_params(value):
if not isinstance(value, dict) or value.keys() - {
"rolloutSteps",
"objectPosition",
"goalPosition",
"stage",
"sourceJobId",
"positionJitter",
"evaluationEpisodes",
"navigationBootstrapSteps",
}:
raise ValueError("mobileParams 包含未知参数")
steps = value.get("rolloutSteps", 128)
if isinstance(steps, bool) or not isinstance(steps, int) or not 8 <= steps <= 4096:
raise ValueError("rolloutSteps 必须为 8–4096 的整数")
stage = value.get("stage", "navigate")
if stage not in ("navigate", "reach", "pick-place"):
raise ValueError("stage 必须为 navigate / reach / pick-place")
source = value.get("sourceJobId")
if source is not None and (
not isinstance(source, str) or not re.fullmatch(r"[0-9a-f]{32}", source)
):
raise ValueError("sourceJobId 必须是服务内的训练作业 ID,不接受 checkpoint 路径")
jitter = value.get("positionJitter", 0.1)
if isinstance(jitter, bool) or not isinstance(jitter, (int, float)) or not 0 <= jitter <= 0.3:
raise ValueError("positionJitter 必须在 0–0.3 m 内")
episodes = value.get("evaluationEpisodes", 10)
if isinstance(episodes, bool) or not isinstance(episodes, int) or not 2 <= episodes <= 64:
raise ValueError("evaluationEpisodes 必须为 2–64 的整数")
bootstrap = value.get("navigationBootstrapSteps", 4096)
if isinstance(bootstrap, bool) or not isinstance(bootstrap, int) or not 0 <= bootstrap <= 10000:
raise ValueError("navigationBootstrapSteps 必须为 0–10000 的整数")
result = {
"navigationBootstrapSteps": bootstrap,
"rolloutSteps": steps,
"stage": stage,
"positionJitter": jitter,
"evaluationEpisodes": episodes,
}
if source is not None:
result["sourceJobId"] = source
for key, default in (
("objectPosition", TASK["objectStart"]),
("goalPosition", TASK["goalStart"]),
):
position = value.get(key, default)
if (
not isinstance(position, list)
or len(position) != 3
or any(
isinstance(x, bool)
or not isinstance(x, (float, int))
or not -TASK["positionScale"] <= x <= TASK["positionScale"]
or not math.isfinite(x)
for x in position
)
or position[2] < TASK["objectStart"][2]
):
raise ValueError(f"{key} 必须是任务范围内的三维坐标,z 不低于物体半高")
result[key] = position.copy()
return result
def mobile_metadata(task_id):
return {
"id": task_id,
"name": f"移动操作 · {ROBOTS[MOBILE_TASKS[task_id]]['label']}",
"family": "mobile-manipulator",
"robotId": MOBILE_TASKS[task_id],
"browserCompatible": True,
"terrainPresets": [],
"terrainParameters": {},
"sensorTypes": [],
"sensorParameters": {},
"mapSyncScope": "当前移动操作场景快照",
"controlDt": TASK["controlDt"],
}
+296
View File
@@ -0,0 +1,296 @@
"""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.
@@ -0,0 +1,52 @@
"""Seeded held-out rollouts; successful export is never a task-success metric."""
from .env import MobileManipulatorEnv
def evaluate_policy(agent, package, params, reset_options, seed):
env = MobileManipulatorEnv(
package,
stage=params["stage"],
reset_options=reset_options,
position_jitter=params["positionJitter"],
)
episodes = []
try:
for i in range(params["evaluationEpisodes"]):
observation, _ = env.reset(seed=seed + i)
done = False
reward_sum = 0.0
while not done:
action, _ = agent.predict(observation, deterministic=True)
observation, reward, terminated, truncated, info = env.step(action)
done = terminated or truncated
reward_sum += reward
episodes.append(
{
"seed": seed + i,
"success": info["is_success"],
"safetyStop": info["safety_stop"],
"steps": env.kernel.steps,
"maxJointVelocity": info["max_joint_velocity"],
"navigationDistance": info["navigation_distance"],
"reward": reward_sum,
}
)
print(
f"Evaluation episode {i + 1}/{params['evaluationEpisodes']}: "
f"success={info['is_success']} safety={info['safety_stop'] or 'none'} "
f"max_joint_velocity={info['max_joint_velocity']:.5f}",
flush=True,
)
finally:
env.close()
return {
"episodes": len(episodes),
"successRate": sum(x["success"] for x in episodes) / len(episodes),
"safetyStops": sum(bool(x["safetyStop"]) for x in episodes),
"maxJointVelocity": max(x["maxJointVelocity"] for x in episodes),
"meanNavigationDistance": sum(x["navigationDistance"] for x in episodes) / len(episodes),
"seed": seed,
"positionJitter": params["positionJitter"],
"rollouts": episodes,
}
@@ -0,0 +1,120 @@
"""Export a trusted PyTorch actor accepting normalized [1,92], returning [1,12].
CLI input is a TorchScript actor, not an entire PPO checkpoint. Any training-time
VecNormalize must be folded into the actor before export. Never load untrusted .pt.
"""
import argparse
import hashlib
import json
from pathlib import Path
import numpy as np
import torch
from .kernel import STAGES, TASK, validate_config
class BoundedActor(torch.nn.Module):
def __init__(self, actor):
super().__init__()
self.actor = actor
def forward(self, observation):
return self.actor(observation).clamp(-1, 1)
def export_policy(actor, package, output, stage="navigate"):
import onnxruntime as ort
if stage not in STAGES:
raise ValueError("invalid training stage")
package, output = Path(package), Path(output)
config_bytes = (package / "robot.json").read_bytes()
config = json.loads(config_bytes)
validate_config(config)
if json.loads((package / "task.json").read_text()) != TASK:
raise ValueError("task contract mismatch")
# Require the compact browser-exported config for the deployment fingerprint.
# Reformatting robot.json changes its hash; export again rather than guessing.
actor = BoundedActor(actor).cpu().eval()
example = torch.zeros((1, TASK["observationSize"]), dtype=torch.float32)
with torch.no_grad():
result = actor(example)
if (
result.shape != (1, TASK["actionSize"])
or result.dtype != torch.float32
or not torch.isfinite(result).all()
):
raise ValueError("actor must return finite float32 [1,12]")
output.parent.mkdir(parents=True, exist_ok=True)
torch.onnx.export(
actor,
example,
str(output),
input_names=["observation"],
output_names=["action"],
opset_version=17,
dynamo=False,
)
session = ort.InferenceSession(str(output), providers=["CPUExecutionProvider"])
rng = np.random.default_rng(42)
for _ in range(5):
obs = rng.uniform(-1, 1, (1, TASK["observationSize"])).astype(np.float32)
actual = session.run(["action"], {"observation": obs})[0]
with torch.no_grad():
expected = actor(torch.from_numpy(obs)).numpy()
np.testing.assert_allclose(actual, expected, atol=1e-5, rtol=1e-5)
metadata = {
"taskId": TASK["id"],
"actionSemantics": TASK["actionSemantics"],
"trainingStage": stage,
"robotId": config["id"],
"observationSize": TASK["observationSize"],
"actionSize": TASK["actionSize"],
"controlDt": TASK["controlDt"],
"normalized": True,
"modelSha256": hashlib.sha256(output.read_bytes()).hexdigest(),
"robotConfigSha256": hashlib.sha256(config_bytes).hexdigest(),
"sceneSha256": hashlib.sha256(
(package / json.loads((package / "environment.json").read_text())["scene"]).read_bytes()
).hexdigest(),
"input": {"name": "observation", "dtype": "float32", "shape": [1, TASK["observationSize"]]},
"output": {"name": "action", "dtype": "float32", "shape": [1, TASK["actionSize"]]},
}
output.with_suffix(".json").write_text(json.dumps(metadata, indent=2) + "\n")
return metadata
class SmokeActor(torch.nn.Module):
"""Untrained hold policy; only validates transport/shape, NOT task competence."""
def __init__(self, config):
super().__init__()
action = torch.zeros((1, TASK["actionSize"]))
self.register_buffer("action", action)
def forward(self, observation):
return observation[:, : TASK["actionSize"]] * 0 + self.action
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--package", required=True, help="extracted browser training bundle")
choice = parser.add_mutually_exclusive_group(required=True)
choice.add_argument("--actor", help="trusted TorchScript actor.pt")
choice.add_argument(
"--smoke", action="store_true", help="UNTRAINED hold policy for wiring tests"
)
parser.add_argument("--output", required=True)
args = parser.parse_args()
actor = (
SmokeActor(json.loads((Path(args.package) / "robot.json").read_text()))
if args.smoke
else torch.jit.load(args.actor, map_location="cpu")
)
export_policy(actor, args.package, args.output)
if __name__ == "__main__":
main()
@@ -0,0 +1,309 @@
"""Math mirror of web_platform/src/mobile/TaskKernel.ts (SI, world, wxyz)."""
import json
import math
from pathlib import Path
import numpy as np
CONTRACTS = Path(__file__).resolve().parents[2] / "contracts"
TASK = json.loads((CONTRACTS / "mobile-manipulator-v2.json").read_text())
ROBOTS = json.loads((CONTRACTS / "mobile-robots-v1.json").read_text())
def clip(value, lo=-1.0, hi=1.0):
return max(lo, min(hi, value))
def rotate(q, v):
w, x, y, z = q
vx, vy, vz = v
tx, ty, tz = 2 * (y * vz - z * vy), 2 * (z * vx - x * vz), 2 * (x * vy - y * vx)
return [
vx + w * tx + y * tz - z * ty,
vy + w * ty + z * tx - x * tz,
vz + w * tz + x * ty - y * tx,
]
def canonical(q):
n = math.hypot(*q)
return np.asarray(q) * ((-1 if q[0] < 0 else 1) / n) if n > 1e-12 else np.array([1, 0, 0, 0])
def relative(parent, child):
w, x, y, z = parent[3:] * np.array([1, -1, -1, -1])
a, b, c, d = child[3:]
q = [
w * a - x * b - y * c - z * d,
w * b + x * a + y * d - z * c,
w * c - x * d + y * a + z * b,
w * d + x * c - y * b + z * a,
]
pos = np.clip(
np.array(rotate([w, x, y, z], child[:3] - parent[:3])) / TASK["positionScale"], -1, 1
)
return np.concatenate((pos, canonical(q)))
def validate_config(c):
primary = [
g for g in c["gripperActuators"] if g.get("joint", c["gripperJoint"]) == c["gripperJoint"]
]
if not primary or any(
g["closed"] != c["gripperClosed"] or g["open"] != c["gripperOpen"] for g in primary
):
raise ValueError("gripper observation/actuator stroke mismatch")
joints = [
*c["baseJoints"],
*(j["name"] for j in c["armJoints"]),
*set(
[c["gripperJoint"], *(g.get("joint", c["gripperJoint"]) for g in c["gripperActuators"])]
),
]
actuators = [
*c["baseActuators"],
*c["armActuators"],
*(g["name"] for g in c["gripperActuators"]),
]
if (
not c["id"]
or c["recipe"] not in ("lekiwi-v1", "lekiwi-bundle", "mjcf")
or not 0 < len(c["armJoints"]) <= TASK["maxArmJoints"]
or len(c["armJoints"]) != len(c["armActuators"])
or not len(c["baseJoints"]) == len(c["baseActuators"]) == len(c["baseMix"])
or not c["baseJoints"]
or not c["gripperActuators"]
or len(set(joints)) != len(joints)
or len(set(actuators)) != len(actuators)
):
raise ValueError("invalid RobotConfig topology")
if (
np.asarray(c["baseMix"]).shape != (len(c["baseJoints"]), 3)
or not np.isfinite(c["baseMix"]).all()
):
raise ValueError("invalid baseMix")
if (
len(c["baseLimits"]) != 3
or not np.isfinite(c["baseLimits"]).all()
or min(c["baseLimits"]) <= 0
or not math.isfinite(c["wheelLimit"])
or c["wheelLimit"] <= 0
or len(c["eefOffset"]) != 3
or not np.isfinite(c["eefOffset"]).all()
):
raise ValueError("invalid scales")
for j in c["armJoints"]:
if (
not np.isfinite([j["min"], j["max"], j["neutral"], j["velocityLimit"]]).all()
or not j["min"] <= j["neutral"] <= j["max"]
or j["min"] >= j["max"]
or j["velocityLimit"] <= 0
or j["mode"] not in ("position", "velocity")
):
raise ValueError("invalid arm limits")
for closed, opened in [
(c["gripperClosed"], c["gripperOpen"]),
*((g["closed"], g["open"]) for g in c["gripperActuators"]),
]:
if not np.isfinite([closed, opened]).all() or closed == opened:
raise ValueError("invalid gripper stroke")
def decode_action(config, action, output=None):
"""Legacy v1 fixture oracle only; v2 environments use SafeActionController."""
action = np.asarray(action, dtype=np.float32)
if action.shape != (TASK["actionSize"],) or not np.isfinite(action).all():
raise ValueError("action must be finite [12]")
nbase, narm = len(config["baseJoints"]), len(config["armJoints"])
if output is None:
output = np.zeros(nbase + narm + len(config["gripperActuators"]), dtype=np.float64)
largest = config["wheelLimit"]
for i, row in enumerate(config["baseMix"]):
output[i] = sum(row[j] * clip(float(action[j])) * config["baseLimits"][j] for j in range(3))
largest = max(largest, abs(output[i]))
output[:nbase] *= config["wheelLimit"] / largest
for i, j in enumerate(config["armJoints"]):
a = clip(float(action[3 + i]))
output[nbase + i] = (
j["min"] + (a + 1) * 0.5 * (j["max"] - j["min"])
if j["mode"] == "position"
else a * j["velocityLimit"]
)
opening = (clip(float(action[11])) + 1) * 0.5
for i, g in enumerate(config["gripperActuators"]):
output[nbase + narm + i] = g["closed"] + opening * (g["open"] - g["closed"])
return output
STAGES = ("navigate", "reach", "pick-place")
def navigation_error(s):
distance = math.hypot(
s[37] - TASK["navigationOffset"][0] - s[0], s[38] - TASK["navigationOffset"][1] - s[1]
)
w, x, y, z = s[3:7]
yaw = math.atan2(2 * (w * z + x * y), 1 - 2 * (y * y + z * z))
return distance, yaw
class TaskKernel:
def __init__(self, config, stage="pick-place"):
validate_config(config)
if stage not in STAGES:
raise ValueError("invalid training stage")
self.stage = stage
self.last_action = np.zeros(TASK["actionSize"], dtype=np.float32)
self.targets = np.zeros(TASK["actionSize"], dtype=np.float32)
self.config = config
self.observation = np.zeros(TASK["observationSize"], dtype=np.float32)
self.reset()
def reset(self):
self.steps = self.settle = 0
self.has_lifted = self.terminated = self.truncated = False
self.reward = self.action_rate = 0.0
self.previous_distance = None
self.last_action.fill(0)
self.targets.fill(0)
self.info = {
"reward_components": dict.fromkeys(
[
"reach",
"lift",
"transport",
"success",
"navigation",
"action_rate",
"joint_velocity",
"safety",
],
0.0,
),
"is_success": False,
"stage": "navigate" if self.stage == "navigate" else "reach",
"safety_stop": "",
"navigation_distance": 0.0,
"max_joint_velocity": 0.0,
}
def record_action(self, applied, targets):
self.action_rate = sum(
(float(applied[i]) - float(self.last_action[i])) ** 2 for i in range(TASK["actionSize"])
)
self.last_action[:] = applied
self.targets[:] = targets
def observe(self, state):
s, o, t = state, self.observation, TASK
o.fill(0)
o[:3] = np.clip(s[:3] / t["positionScale"], -1, 1)
o[3:7] = canonical(s[3:7])
o[7:10] = np.clip(s[7:10] / t["linearVelocityScale"], -1, 1)
o[10:13] = np.clip(s[10:13] / t["angularVelocityScale"], -1, 1)
for i, j in enumerate(self.config["armJoints"]):
o[13 + i] = clip(2 * (s[13 + i] - j["min"]) / (j["max"] - j["min"]) - 1)
o[21 + i] = clip(s[21 + i] / t["jointVelocityScale"])
o[29 + i] = 1
o[37] = clip(2 * s[29] - 1)
o[38:41] = np.clip(s[30:33] / t["positionScale"], -1, 1)
o[41:45] = canonical(s[33:37])
o[45:52] = relative(s[30:37], s[37:44])
o[52:59] = relative(s[47:54], s[37:44])
o[59:62] = np.clip(s[47:50] / t["positionScale"], -1, 1)
o[62:66] = canonical(s[50:54])
o[66], o[67] = self.has_lifted, self.settle / t["settleSteps"]
o[68:80] = self.last_action
o[80:92] = self.targets
return o
def evaluate(self, s, safety_stop="", peak_velocity=0.0):
if self.terminated or self.truncated:
raise RuntimeError("episode ended; reset required")
t = TASK
reach = math.hypot(*(s[37:40] - s[30:33]))
goal = math.hypot(*(s[37:40] - s[47:50]))
lift = clip((s[39] - t["objectStart"][2]) / t["liftHeight"], 0, 1)
if self.stage == "pick-place" and lift >= 1 and reach < t["graspDistance"] and s[29] < 0.4:
self.has_lifted = True
distance, yaw = navigation_error(s)
near = distance < t["navigationTolerance"] and abs(yaw) < t["navigationYawTolerance"]
stopped = (
math.hypot(*s[7:10]) < t["navigationSpeedTolerance"] and math.hypot(*s[10:13]) < 0.1
)
settled = (
self.has_lifted
and goal < t["goalTolerance"]
and s[29] > t["releaseOpening"]
and reach > t["graspDistance"]
and math.hypot(*s[44:47]) < t["settleSpeed"]
)
if self.stage == "navigate":
settled = near and stopped
elif self.stage == "reach":
settled = (
near and stopped and reach < t["graspDistance"] and math.hypot(*s[21:29]) < 0.15
)
self.settle = min(t["settleSteps"], self.settle + 1) if settled else 0
success = self.settle == t["settleSteps"]
r = self.info["reward_components"]
r["reach"] = t["controlDt"] * t["reachWeight"] * math.exp(-t["reachGain"] * reach)
r["lift"] = t["controlDt"] * t["liftWeight"] * lift
r["transport"] = (
t["controlDt"] * t["transportWeight"] * math.exp(-t["transportGain"] * goal)
if self.has_lifted
else 0.0
)
if self.stage == "navigate":
r["reach"] = r["lift"] = r["transport"] = 0.0
elif self.stage == "reach":
r["lift"] = r["transport"] = 0.0
progress = 0 if self.previous_distance is None else self.previous_distance - distance
self.previous_distance = distance
r["navigation"] = (
(
t["navigationProgressWeight"] * progress
- t["controlDt"] * (distance + 0.1 * abs(yaw))
)
if not self.has_lifted
else 0.0
)
r["action_rate"] = -t["actionRateWeight"] * self.action_rate
r["joint_velocity"] = (
-t["controlDt"] * t["jointVelocityWeight"] * sum(float(v) ** 2 for v in s[21:29])
)
peak_velocity = max(peak_velocity, max(abs(s[21:29])))
if not safety_stop and peak_velocity > t["jointSpeedStop"]:
safety_stop = "joint_velocity"
if not safety_stop and (
1 - 2 * (s[4] ** 2 + s[5] ** 2) < 0.5 or max(abs(s[:2])) > t["positionScale"]
):
safety_stop = "base_pose"
if safety_stop:
success = False
r["success"] = t["successBonus"] if success else 0.0
r["safety"] = -t["safetyPenalty"] if safety_stop else 0.0
self.reward = sum(r.values())
self.terminated = success or bool(safety_stop)
self.info["safety_stop"] = safety_stop
self.info["navigation_distance"] = distance
self.info["max_joint_velocity"] = max(self.info["max_joint_velocity"], peak_velocity)
self.steps += 1
self.truncated = self.steps >= t["maxSteps"] and not success
self.info["is_success"] = success
self.info["stage"] = (
"safety-stop"
if safety_stop
else "success"
if success
else "navigate"
if self.stage == "navigate" or (not near and not self.has_lifted)
else "transport"
if self.has_lifted
else "lift"
if reach < t["graspDistance"]
else "reach"
)
self.observe(s)
return self.observation, self.reward, self.terminated, self.truncated, self.info
@@ -0,0 +1,118 @@
"""Versioned, stateful action adapter mirrored by mobile/SafeActionController.ts.
Position targets are integrated at bounded speed, not remapped over full joint travel.
Both applied actions and target state are observable; no hidden policy-side filter.
"""
import numpy as np
from .kernel import TASK as T
from .kernel import clip, navigation_error
class SafeActionController:
def __init__(self, config):
self.config = config
self.applied = np.zeros(T["actionSize"], dtype=np.float32)
self.targets = np.zeros(T["actionSize"], dtype=np.float32)
self.control = np.zeros(
len(config["baseJoints"]) + len(config["armJoints"]) + len(config["gripperActuators"])
)
def reset(self, state, previous_control=None):
previous_control = previous_control.copy() if previous_control is not None else None
self.applied.fill(0)
self.targets.fill(0)
self.control.fill(0)
n = len(self.config["baseJoints"])
for i, j in enumerate(self.config["armJoints"]):
if j["mode"] == "position":
q = float(state[13 + i])
if previous_control is not None and np.isfinite(previous_control[n + i]):
q = clip(
previous_control[n + i],
q - T["armTrackingError"],
q + T["armTrackingError"],
)
q = clip(q, j["min"], j["max"])
self.control[n + i] = q
self.targets[3 + i] = 2 * (q - j["min"]) / (j["max"] - j["min"]) - 1
opening = float(state[29])
for i, g in enumerate(self.config["gripperActuators"]):
index = n + len(self.config["armJoints"]) + i
if (
previous_control is not None
and g.get("joint", self.config["gripperJoint"]) == self.config["gripperJoint"]
and np.isfinite(previous_control[index])
):
opening = clip(
(previous_control[index] - g["closed"]) / (g["open"] - g["closed"]), 0, 1
)
break
self.targets[11] = 2 * opening - 1
self._gripper(opening)
def _gripper(self, opening):
n = len(self.config["baseJoints"]) + len(self.config["armJoints"])
for i, g in enumerate(self.config["gripperActuators"]):
self.control[n + i] = g["closed"] + opening * (g["open"] - g["closed"])
def apply(self, action, state, stage, lifted):
action = np.asarray(action, dtype=np.float32)
if action.shape != (T["actionSize"],) or not np.isfinite(action).all():
raise ValueError("action must be finite [12]")
distance, yaw = navigation_error(state)
near = distance < T["navigationTolerance"] and abs(yaw) < T["navigationYawTolerance"]
manipulate = stage != "navigate" and (near or lifted)
c, dt = self.config, T["controlDt"]
# Actual body-frame velocity commands have bounded acceleration.
for i in range(3):
limit = min(c["baseLimits"][i], T["baseSpeedLimits"][i])
desired = 0 if near and not lifted else clip(float(action[i]))
delta = T["baseAccelerationLimits"][i] * dt / limit
self.applied[i] = clip(desired, self.applied[i] - delta, self.applied[i] + delta)
largest = c["wheelLimit"]
for i, row in enumerate(c["baseMix"]):
self.control[i] = sum(
row[j] * self.applied[j] * min(c["baseLimits"][j], T["baseSpeedLimits"][j])
for j in range(3)
)
largest = max(largest, abs(self.control[i]))
n = len(c["baseJoints"])
self.control[:n] *= c["wheelLimit"] / largest
for i, j in enumerate(c["armJoints"]):
k = 3 + i
speed = min(j["velocityLimit"], T["armSpeedLimit"])
delta = T["armAccelerationLimit"] * dt / speed
a = clip(float(action[k])) if manipulate else 0.0
a = clip(a, self.applied[k] - delta, self.applied[k] + delta)
if j["mode"] == "position":
old = self.control[n + i]
desired = clip(
old + a * speed * dt,
state[13 + i] - T["armTrackingError"],
state[13 + i] + T["armTrackingError"],
)
target = clip(clip(desired, old - speed * dt, old + speed * dt), j["min"], j["max"])
# During navigation hold the reset/current target, not the gravity-sagged qpos.
if not manipulate:
target = old
self.control[n + i] = target
self.applied[k] = (target - old) / (speed * dt)
self.targets[k] = 2 * (target - j["min"]) / (j["max"] - j["min"]) - 1
else:
self.applied[k] = a if manipulate else 0
self.control[n + i] = self.applied[k] * speed
opening = (float(self.targets[11]) + 1) / 2
target = clip(
opening
+ (clip(float(action[11])) if manipulate and stage == "pick-place" else 0)
* T["gripperOpeningRate"]
* dt,
0,
1,
)
self.applied[11] = (target - opening) / (T["gripperOpeningRate"] * dt)
self.targets[11] = 2 * target - 1
self._gripper(target)
return self.control
@@ -0,0 +1,161 @@
"""Authenticated browser scene snapshots; no caller-supplied filesystem paths."""
import hashlib
import json
import re
import shutil
import stat
import tempfile
import threading
import zipfile
from pathlib import Path, PurePosixPath
from xml.etree import ElementTree as ET
from .config import ROBOTS, TASK
MAX_UPLOAD = 128 * 1024**2
MAX_EXPANDED = 512 * 1024**2
class MobilePackages:
def __init__(self, root):
self.root = Path(root)
self.lock = threading.Lock()
def receive(self, stream, length):
if not 0 < length <= MAX_UPLOAD:
raise ValueError("场景上传大小必须在 1–128 MiB 内")
self.root.mkdir(parents=True, exist_ok=True)
with self.lock, tempfile.TemporaryDirectory(dir=self.root) as temporary:
temp = Path(temporary)
archive = temp / "upload.zip"
digest = hashlib.sha256()
with archive.open("wb") as out:
remaining = length
while remaining:
chunk = stream.read(min(1024 * 1024, remaining))
if not chunk:
raise ValueError("场景上传不完整")
digest.update(chunk)
out.write(chunk)
remaining -= len(chunk)
package_id = digest.hexdigest()
destination = self.root / package_id
if destination.is_dir():
return {"id": package_id, **self.describe(package_id)}
if sum(p.is_dir() and len(p.name) == 64 for p in self.root.iterdir()) >= 20:
raise ValueError("场景快照已达 20 份,请在停止服务后清理 mobile_packages")
directory = temp / "package"
directory.mkdir()
try:
with zipfile.ZipFile(archive) as z:
infos = z.infolist()
if len(infos) > 10000 or sum(i.file_size for i in infos) > MAX_EXPANDED:
raise ValueError("场景展开超出 512 MiB / 10000 文件上限")
names = set()
for info in infos:
name = info.filename
path = PurePosixPath(name)
if (
not name
or "\\" in name
or ":" in name
or path.is_absolute()
or ".." in path.parts
or str(path) in names
or stat.S_ISLNK(info.external_attr >> 16)
):
raise ValueError("场景包含不安全或重复路径")
names.add(str(path))
if not info.is_dir():
target = directory.joinpath(*path.parts)
target.parent.mkdir(parents=True, exist_ok=True)
with z.open(info) as src, target.open("wb") as dst:
shutil.copyfileobj(src, dst)
self._validate(directory)
except (
zipfile.BadZipFile,
KeyError,
TypeError,
ET.ParseError,
OSError,
NotImplementedError,
RuntimeError,
) as error:
raise ValueError(f"场景快照无效:{error}") from error
directory.rename(destination)
return {"id": package_id, **self.describe(package_id)}
def path(self, package_id):
if not isinstance(package_id, str) or not re.fullmatch(r"[0-9a-f]{64}", package_id):
raise ValueError("mobilePackageId 无效")
directory = self.root / package_id
if not directory.is_dir():
raise ValueError("移动操作场景不存在,请重新开始训练以同步场景")
return directory
def describe(self, package_id):
return self._validate(self.path(package_id))
@staticmethod
def _validate(directory):
for name in ("environment.json", "robot.json", "task.json"):
if (directory / name).stat().st_size > 128 * 1024:
raise ValueError("场景契约 JSON 超过 128 KiB")
metadata = json.loads((directory / "environment.json").read_text())
robot_bytes = (directory / "robot.json").read_bytes()
robot = json.loads(robot_bytes)
if (
not isinstance(metadata, dict)
or not isinstance(robot, dict)
or not isinstance(robot.get("id"), str)
or robot.get("id") not in ROBOTS
or robot != ROBOTS[robot["id"]]
or json.loads((directory / "task.json").read_text()) != TASK
or metadata.get("robotId") != robot["id"]
or metadata.get("taskId") != TASK["id"]
):
raise ValueError("场景机器人/任务契约与已注册变体不匹配")
if metadata.get("mujoco") != "3.11.0":
raise ValueError("场景需要与浏览器一致的 MuJoCo 3.11.0")
scene_name = metadata.get("scene")
if not isinstance(scene_name, str) or "\\" in scene_name or ":" in scene_name:
raise ValueError("scene 路径无效")
scene = (directory / scene_name).resolve()
if not scene.is_relative_to(directory.resolve()) or not scene.is_file():
raise ValueError("scene 越界或不存在")
if scene.stat().st_size > 16 * 1024**2:
raise ValueError("场景 MJCF 超过 16 MiB")
data = scene.read_bytes()
if b"<!DOCTYPE" in data.upper() or b"<!ENTITY" in data.upper():
raise ValueError("场景不接受 DTD/entity")
root = ET.fromstring(data)
if (
root.tag != "mujoco"
or root.find(".//include") is not None
or root.find(".//plugin") is not None
):
raise ValueError("必须上传展开后的 MJCF,不允许 include/plugin")
compiler = root.find("compiler")
for element in root.iter():
for key, value in element.attrib.items():
# Includes texture cubemap fileleft/fileright/fileup/... references.
is_file = key.startswith("file")
if not is_file and key not in ("meshdir", "texturedir", "assetdir"):
continue
if "\\" in value or ":" in value or PurePosixPath(value).is_absolute():
raise ValueError("场景资产必须是包内相对路径")
prefix = ""
if is_file and compiler is not None:
prefix = compiler.get("meshdir" if element.tag == "mesh" else "texturedir", "")
prefix = prefix or compiler.get("assetdir", "")
target = (scene.parent / prefix / value).resolve()
if not target.is_relative_to(directory.resolve()):
raise ValueError("场景资产引用越界")
if is_file and not target.is_file():
raise ValueError(f"场景缺少资产:{value}")
return {
"robotId": robot["id"],
"sceneSha256": hashlib.sha256(data).hexdigest(),
"robotConfigSha256": hashlib.sha256(robot_bytes).hexdigest(),
}
@@ -0,0 +1,8 @@
# MuJoCo must match environment.json from the browser (currently WASM 3.11.0).
mujoco==3.11.0
gymnasium>=1.0,<2
stable-baselines3>=2.6,<3
numpy>=2.0
torch>=2.6
onnx>=1.17
onnxruntime>=1.20
+179
View File
@@ -0,0 +1,179 @@
"""Server-owned SB3 PPO runner. stdout follows the shared iteration/scalar protocol."""
import argparse
import json
from pathlib import Path
import numpy as np
import torch
from .bootstrap import bootstrap_navigation
from .config import MOBILE_TASKS, validate_mobile_params
from .env import MobileManipulatorEnv
from .evaluation import evaluate_policy
from .export_onnx import export_policy
class DeterministicActor(torch.nn.Module):
def __init__(self, policy):
super().__init__()
self.policy = policy
def forward(self, observation):
return self.policy(observation, deterministic=True)[0]
def main():
from stable_baselines3 import PPO
from stable_baselines3.common.monitor import Monitor
from stable_baselines3.common.vec_env import DummyVecEnv
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--package", required=True)
parser.add_argument("--iterations", type=int, default=1000)
parser.add_argument("--num-envs", type=int, default=1)
parser.add_argument("--output", required=True)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--device", default="cpu")
parser.add_argument("--task-id", required=True, choices=MOBILE_TASKS)
parser.add_argument("--params", default="{}")
parser.add_argument("--resume", help="server-owned trusted PPO checkpoint only")
args = parser.parse_args()
torch.set_num_threads(min(4, torch.get_num_threads()))
params = validate_mobile_params(json.loads(args.params))
if args.iterations < 1 or not 1 <= args.num_envs <= 64:
parser.error("iterations/num-envs out of range")
if args.device.startswith("cuda") and not torch.cuda.is_available():
raise RuntimeError("选择了 GPU,但训练 Python 中 CUDA 不可用;请改用 CPU")
if params["stage"] != "navigate" and not args.resume:
parser.error("later stages require a validated previous-stage checkpoint")
reset_options = {"object": params["objectPosition"], "goal": params["goalPosition"]}
print(f"Training stage: {params['stage']} | rate-limited-position-target-v1", flush=True)
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
class ConsolePPO(PPO):
iteration = 0
def train(self):
super().train()
self.iteration += 1
print(f"Learning iteration {self.iteration} / {args.iterations}", flush=True)
print(f"Total timesteps: {self.num_timesteps}", flush=True)
for key, label in (
("value_loss", "value loss"),
("policy_gradient_loss", "surrogate loss"),
("entropy_loss", "entropy loss"),
):
value = self.logger.name_to_value.get(f"train/{key}")
if value is not None and np.isfinite(value):
print(f"Mean {label}: {value:.9g}", flush=True)
print(f"Mean reward: {np.mean(self.rollout_buffer.rewards):.9g}", flush=True)
if self.ep_info_buffer:
success = np.mean([x["is_success"] for x in self.ep_info_buffer])
peak = max(x["max_joint_velocity"] for x in self.ep_info_buffer)
print(f"Mean success rate: {success:.6g}", flush=True)
print(f"Max joint velocity: {peak:.6g}", flush=True)
print(
f"Mean episode length: {np.mean([x['l'] for x in self.ep_info_buffer]):.9g}",
flush=True,
)
def make_env():
return Monitor(
MobileManipulatorEnv(
args.package,
reset_options=reset_options,
stage=params["stage"],
position_jitter=params["positionJitter"],
),
info_keywords=("is_success", "max_joint_velocity", "safety_stop"),
)
# Native MuJoCo CPU simulation, vectorized rollout; device selects the PPO network.
env = DummyVecEnv([make_env for _ in range(args.num_envs)])
try:
if env.envs[0].unwrapped.config["id"] != MOBILE_TASKS[args.task_id]:
raise ValueError("task/robot mismatch")
rollout = params["rolloutSteps"] * args.num_envs
batch = min(64, rollout)
while rollout % batch:
batch -= 1
settings = dict(
seed=args.seed,
verbose=0,
device=args.device,
n_steps=params["rolloutSteps"],
batch_size=batch,
)
if args.resume:
agent = ConsolePPO.load(args.resume, env=env, **settings)
agent.iteration = 0
if agent.ep_info_buffer is not None:
agent.ep_info_buffer.clear()
if getattr(agent, "training_stage", params["stage"]) != params["stage"]:
# Explore the newly enabled arm without destroying base navigation.
with torch.no_grad():
agent.policy.log_std[3:].fill_(-1.5)
else:
agent = ConsolePPO(
"MlpPolicy",
env,
**settings,
policy_kwargs=dict(
net_arch=[128, 128],
activation_fn=torch.nn.ELU,
log_std_init=-3.0 if params["navigationBootstrapSteps"] else -1.0,
),
# Conservative PPO updates: do not destroy navigation initialization
# while the critic is still learning the sparse terminal return.
learning_rate=3e-5,
n_epochs=4,
ent_coef=0.001,
target_kl=0.005,
)
if not args.resume:
bootstrap_navigation(agent, args.package, params, reset_options, args.seed + 50_000)
agent.learn(
total_timesteps=args.iterations * rollout,
log_interval=None,
reset_num_timesteps=not bool(args.resume),
)
agent.training_stage = params["stage"]
agent.save(str(output.with_suffix(".ppo.zip")))
# Distinct seeds, deterministic actions and the same controller used by deployment.
evaluation = evaluate_policy(
agent, args.package, params, reset_options, args.seed + 100_000
)
(output.parent / "evaluation.json").write_text(json.dumps(evaluation, indent=2) + "\n")
metadata = export_policy(
DeterministicActor(agent.policy), args.package, output, params["stage"]
)
metadata.update(
version=1,
browserCompatible=True,
trainingTaskId=args.task_id,
seed=args.seed,
resetOptions=reset_options,
trainingParams=params,
trainedTimesteps=agent.num_timesteps,
initialization="resumed-checkpoint"
if args.resume
else "navigation-BC-then-PPO"
if params["navigationBootstrapSteps"]
else "PPO-from-scratch",
evaluation={k: v for k, v in evaluation.items() if k != "rollouts"},
)
output.with_name("deployment.json").write_text(json.dumps(metadata, indent=2) + "\n")
print(
f"Evaluation success rate: {evaluation['successRate']:.3f}; "
f"safety stops: {evaluation['safetyStops']}",
flush=True,
)
print("ONNX export complete (export does NOT imply task competence)", flush=True)
finally:
env.close()
if __name__ == "__main__":
main()
@@ -0,0 +1,47 @@
"""Compare browser-exported fixed actions/rollout with the same native MJCF."""
import argparse
import json
from pathlib import Path
import mujoco
import numpy as np
from gymnasium.utils.env_checker import check_env
from .env import MobileManipulatorEnv
def validate(package, rollout, atol=2e-5):
reference = json.loads(Path(rollout).read_text())
env = MobileManipulatorEnv(package)
maxima = dict(qpos=0.0, observation=0.0, control=0.0, reward=0.0)
try:
for action, expected in zip(reference["actions"], reference["rollout"], strict=True):
observation, reward, _, _, info = env.step(action)
for name, actual, target in [
("qpos", env.data.qpos, expected["qpos"]),
("observation", observation, expected["observation"]),
("control", env.data.ctrl, expected["ctrl"]),
("reward", reward, expected["reward"]),
]:
error = float(np.max(np.abs(np.asarray(actual) - np.asarray(target))))
maxima[name] = max(maxima[name], error)
np.testing.assert_allclose(actual, target, atol=atol, rtol=0, err_msg=name)
if info["stage"] != expected["stage"]:
raise AssertionError("stage mismatch")
check_env(env, skip_render_check=True)
finally:
env.close()
return dict(mujoco=mujoco.__version__, steps=len(reference["actions"]), max_error=maxima)
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--package", required=True)
parser.add_argument("--rollout", required=True)
args = parser.parse_args()
print(json.dumps(validate(args.package, args.rollout), indent=2))
if __name__ == "__main__":
main()
+273 -13
View File
@@ -25,6 +25,9 @@ from pathlib import Path
from typing import Any
from urllib.parse import parse_qs, unquote, urlsplit
from mobile_manipulator.config import MOBILE_TASKS, validate_mobile_params
from mobile_manipulator.config import TASK as MOBILE_CONTRACT
from mobile_manipulator.packages import MAX_UPLOAD, MobilePackages
from pretrained_sources import PretrainedSources, SourceError
from task_config import (
OBSTACLE_TASK,
@@ -38,10 +41,10 @@ from tuning.process import GpuLease, ResourceBusyError
from tuning.schema import RewardConfigError, validate_configuration
from tuning.scoring import EvaluationError
VERSION = "0.4.0"
VERSION = "0.6.0"
MAX_REQUEST_BYTES = 128 * 1024 # Bounded full boxes-v1 payload (<=257 boxes).
# Rough 可以训练,但其高度扫描 actor 不允许冒充浏览器 Flat 部署。
DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK)
DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK, *MOBILE_TASKS)
ACTIVE_STATES = {"queued", "running"}
MAX_JOBS = 20
ANSI_ESCAPE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]")
@@ -86,6 +89,9 @@ class TrainingConfig:
task_config: dict[str, Any] | None = None
deployment: dict[str, Any] = field(default_factory=dict)
pretrained: dict[str, Any] | None = None
mobile_package_id: str | None = None
mobile_params: dict[str, Any] | None = None
mobile_checkpoint: str | None = None
@dataclass
@@ -135,9 +141,11 @@ class TrainingManager:
check_environment: bool = True,
lease: GpuLease | None = None,
sources: PretrainedSources | None = None,
mobile_python: str | None = None,
):
self.trainer_root = trainer_root.expanduser().resolve()
self.python = str(Path(python).expanduser()) if os.sep in python else python
self.mobile_python = mobile_python or self.python
self.tasks = tasks
self.jobs: dict[str, TrainingJob] = {}
self.lock = threading.RLock()
@@ -146,8 +154,39 @@ class TrainingManager:
self.lease = lease or GpuLease()
self.preset_resolver: Any = None
self.sources = sources
self.mobile_packages = MobilePackages(self.trainer_root / "logs" / "mobile_packages")
self._mobile_environment_error: str | None | bool = False
def readiness_error(self) -> str | None:
def readiness_error(self, task_id: str | None = None) -> str | None:
if task_id in MOBILE_TASKS:
if not self.check_environment:
return None
if self._mobile_environment_error is False:
try:
result = subprocess.run(
[
self.mobile_python,
"-c",
(
"import mujoco,gymnasium,torch,stable_baselines3,onnx,onnxruntime; "
"assert mujoco.__version__ == '3.11.0', "
"'移动操作需要 MuJoCo 3.11.0,请配置 --mobile-python;"
"不要升级 Go2 环境'"
),
],
capture_output=True,
text=True,
timeout=30,
check=False,
)
self._mobile_environment_error = (
"移动操作 Python 依赖不可用:" + result.stderr[-1500:]
if result.returncode
else None
)
except (OSError, subprocess.TimeoutExpired) as error:
self._mobile_environment_error = f"无法检查移动操作环境:{error}"
return self._mobile_environment_error or None
if not self.trainer_root.is_dir():
return f"训练工程目录不存在:{self.trainer_root}"
if not (self.trainer_root / "scripts" / "train.py").is_file():
@@ -184,20 +223,26 @@ class TrainingManager:
return next((job.id for job in self.jobs.values() if job.state in ACTIVE_STATES), None)
def health(self) -> dict[str, Any]:
error = self.readiness_error()
errors = {task: self.readiness_error(task) for task in self.tasks}
ready = any(error is None for error in errors.values())
error = None if ready else next(iter(errors.values()), "没有可用任务")
metadata = task_metadata(self.tasks)
for item in metadata:
item.update(ready=errors[item["id"]] is None, error=errors[item["id"]])
return {
"version": VERSION,
"ready": error is None,
"ready": ready,
"trainerRoot": str(self.trainer_root),
"python": self.python,
"tasks": list(self.tasks),
"pretrainedSources": self.sources.catalog() if self.sources else [],
"pretrainedUpload": {
"enabled": self.sources is not None, "templateId": "go2-legacy47-v1",
"enabled": self.sources is not None,
"templateId": "go2-legacy47-v1",
"formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2},
"endpoint": "/api/training/pretrained-sources/upload",
},
"taskMetadata": task_metadata(self.tasks),
"taskMetadata": metadata,
"activeJobId": self.active_job_id(),
"error": error,
}
@@ -221,6 +266,8 @@ class TrainingManager:
"sensorCfg",
"sensorType",
"customTerrainBoxes",
"mobilePackageId",
"mobileParams",
}
if payload.keys() - allowed:
raise ApiError(HTTPStatus.BAD_REQUEST, "请求包含未知字段(不接受配置路径/MJCF)")
@@ -276,6 +323,56 @@ class TrainingManager:
except RewardConfigError as error:
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
seed = integer("seed", 0, 2_147_483_647)
if task_id in MOBILE_TASKS:
if any(
k in payload
for k in (
"terrainPreset",
"terrainParams",
"sensorCfg",
"sensorType",
"customTerrainBoxes",
"pretrainedSourceId",
)
):
raise ApiError(
HTTPStatus.BAD_REQUEST, "移动操作任务不接受 Go2 地形/传感器/预训练参数"
)
try:
package_id = payload.get("mobilePackageId")
package = self.mobile_packages.describe(package_id)
if package["robotId"] != MOBILE_TASKS[task_id]:
raise ValueError("训练任务与场景机器人变体不匹配")
params = validate_mobile_params(payload.get("mobileParams", {}))
checkpoint = self.mobile_resume(params, task_id, package, package_id)
except (ValueError, OSError) as error:
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
if device == "gpu" and len(raw_gpu_ids) != 1:
raise ApiError(HTTPStatus.BAD_REQUEST, "移动操作 PPO 仅支持单个 GPU")
return TrainingConfig(
task_id=task_id,
num_envs=integer("numEnvs", 1, 64),
max_iterations=integer("maxIterations", 1, 1_000_000),
seed=seed,
run_name=run_name,
device=device,
gpu_ids=raw_gpu_ids,
wandb_mode=wandb_mode,
mobile_package_id=package_id,
mobile_params=params,
mobile_checkpoint=checkpoint,
deployment={
"trainingStage": params["stage"],
"actionSemantics": MOBILE_CONTRACT["actionSemantics"],
"version": 1,
"browserCompatible": True,
"trainingTaskId": task_id,
"taskId": MOBILE_CONTRACT["id"],
**package,
},
)
if "mobilePackageId" in payload or "mobileParams" in payload:
raise ApiError(HTTPStatus.BAD_REQUEST, "Go2 任务不接受移动操作参数")
try:
custom = validate_task_config(task_id, payload, seed)
except TaskConfigError as error:
@@ -308,11 +405,70 @@ class TrainingManager:
reward_config=reward_config,
)
def mobile_resume(self, params, task_id, package, package_id) -> str | None:
source_id = params.get("sourceJobId")
if source_id is None:
if params["stage"] != "navigate":
raise ValueError("请先完成底盘接近训练,再选择通过评估的前一阶段作业")
return None
with self.lock:
source = self.jobs.get(source_id)
if (
not source
or source.state != "succeeded"
or source.config.task_id != task_id
or not source.artifact
):
raise ValueError("接续训练需要同一机器人已完成的服务内作业")
if source.config.mobile_package_id != package_id:
raise ValueError("接续场景/资产快照不匹配;资产变化后请重新训练")
deployment = source.config.deployment
if (
any(
deployment.get(key) != package.get(key)
for key in ("robotId", "sceneSha256", "robotConfigSha256")
)
or deployment.get("actionSemantics") != MOBILE_CONTRACT["actionSemantics"]
or deployment.get("taskId") != MOBILE_CONTRACT["id"]
):
raise ValueError("接续作业与当前场景/安全控制契约不匹配")
stages = ("navigate", "reach", "pick-place")
previous = deployment.get("trainingStage")
if (
previous not in stages
or not stages.index(previous)
<= stages.index(params["stage"])
<= stages.index(previous) + 1
):
raise ValueError("只支持同阶段续训或依次推进:底盘接近 → 末端接近 → 抓取放置")
if previous != params["stage"]:
if any(
params[key] != (source.config.mobile_params or {}).get(key)
for key in ("objectPosition", "goalPosition", "positionJitter")
):
raise ValueError(
"升级阶段必须保留已评估的初态分布;改变坐标/随机范围请先同阶段续训"
)
evaluation = deployment.get("evaluation", {})
if (
evaluation.get("episodes", 0) < 10
or evaluation.get("successRate", 0) < MOBILE_CONTRACT["navigationSuccessRate"]
or evaluation.get("safetyStops", 1) != 0
):
raise ValueError(
"上一阶段尚未达标:至少 10 回合独立评估、成功率 ≥80%、"
"无安全终止;请先同阶段续训"
)
checkpoint = source.artifact.with_suffix(".ppo.zip")
if not checkpoint.is_file():
raise ValueError("接续作业缺少 PPO checkpoint")
return str(checkpoint)
def start(self, payload: Any) -> dict[str, Any]:
error = self.readiness_error()
config = self.parse_config(payload)
error = self.readiness_error(config.task_id)
if error:
raise ApiError(HTTPStatus.SERVICE_UNAVAILABLE, error)
config = self.parse_config(payload)
with self.lock:
if self.active_job_id():
raise ApiError(HTTPStatus.CONFLICT, "已有训练任务正在运行,请等待完成或先停止任务")
@@ -357,6 +513,12 @@ class TrainingManager:
raise ApiError(HTTPStatus.NOT_FOUND, "该训练任务尚未生成 policy.onnx")
return job.artifact
def deployment_artifact(self, job_id: str) -> Path:
artifact = self.artifact(job_id).with_name("deployment.json")
if not artifact.is_file():
raise ApiError(HTTPStatus.NOT_FOUND, "该任务没有独立部署元数据")
return artifact
def cancel(self, job_id: str) -> dict[str, Any]:
with self.lock:
job = self.jobs.get(job_id)
@@ -400,6 +562,27 @@ class TrainingManager:
def command_for(
self, config: TrainingConfig, task_config_path: Path | None = None
) -> list[str]:
if config.task_id in MOBILE_TASKS:
return [
self.mobile_python,
"-u",
"-m",
"training_server.mobile_manipulator.train",
"--package",
str(self.mobile_packages.path(config.mobile_package_id)),
"--iterations",
str(config.max_iterations),
"--num-envs",
str(config.num_envs),
"--seed",
str(config.seed),
"--device",
"cpu" if config.device == "cpu" else f"cuda:{config.gpu_ids[0]}",
"--params",
json.dumps(config.mobile_params),
"--task-id",
config.task_id,
] + (["--resume", config.mobile_checkpoint] if config.mobile_checkpoint else [])
command = [
self.python,
"-u",
@@ -480,6 +663,26 @@ class TrainingManager:
config_path = job_dir / "training_config.json"
config_path.write_text(json.dumps(job.config.task_config), encoding="utf-8")
command = self.command_for(job.config, config_path)
mobile = job.config.task_id in MOBILE_TASKS
if mobile:
job_dir = self.trainer_root / "logs" / "rsl_rl" / "web_jobs" / job.id
job_dir.mkdir(parents=True, exist_ok=True)
command.extend(("--output", str(job_dir / "policy.onnx")))
(job_dir / "training_config.json").write_text(
json.dumps(
{
"taskId": job.config.task_id,
"mobilePackageId": job.config.mobile_package_id,
"mobileParams": job.config.mobile_params,
"numEnvs": job.config.num_envs,
"maxIterations": job.config.max_iterations,
"seed": job.config.seed,
"device": job.config.device,
"gpuIds": job.config.gpu_ids,
}
),
encoding="utf-8",
)
if config_path is not None:
command.extend(("--output-dir", str(config_path.parent)))
# Popen 与 process 登记必须和取消检查处于同一个临界区:cancel() 要么在
@@ -490,7 +693,7 @@ class TrainingManager:
return
process = subprocess.Popen(
command,
cwd=self.trainer_root,
cwd=Path(__file__).resolve().parent.parent if mobile else self.trainer_root,
env=environment,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
@@ -513,7 +716,22 @@ class TrainingManager:
finally:
process.stdout.close()
return_code = process.wait()
artifact = self._find_artifact(before)
artifact = (job_dir / "policy.onnx") if mobile else self._find_artifact(before)
if artifact is not None and not artifact.is_file():
artifact = None
if mobile and return_code == 0 and not job.cancel_requested and artifact:
deployment = json.loads((job_dir / "deployment.json").read_text())
for key in (
"trainingTaskId",
"robotId",
"sceneSha256",
"robotConfigSha256",
"trainingStage",
"actionSemantics",
):
if deployment.get(key) != job.config.deployment.get(key):
raise ValueError(f"导出部署元数据不匹配:{key}")
job.config.deployment = deployment
with self.lock:
job.process = None
job.ended_at = now_iso()
@@ -543,7 +761,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
tuning_manager: TuningManager
allowed_origins: tuple[str, ...] = ()
access_token = ""
server_version = "MuJoCoLocalTraining/0.4"
server_version = "MuJoCoLocalTraining/0.6"
def log_message(self, format: str, *args: Any) -> None:
sys.stderr.write(f"[{self.log_date_time_string()}] {format % args}\n")
@@ -644,7 +862,12 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
raise ApiError(HTTPStatus.BAD_REQUEST, "上传显示名称过长")
try:
result = self.manager.sources.receive_upload(
self.rfile, length, fmt, template, name, set_timeout=self.connection.settimeout,
self.rfile,
length,
fmt,
template,
name,
set_timeout=self.connection.settimeout,
)
except OSError as error:
raise ApiError(
@@ -734,6 +957,14 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
if match:
self._send_file(self.tuning_manager.best_artifact(match.group(1)), "policy.onnx")
return
deployment_match = re.fullmatch(
r"/api/training/jobs/([0-9a-f]{32})/artifacts/deployment\.json", path
)
if deployment_match:
self._send_file(
self.manager.deployment_artifact(deployment_match.group(1)), "deployment.json"
)
return
job_id, artifact = self._route(path)
if not job_id:
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
@@ -751,6 +982,29 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
if path == "/api/training/pretrained-sources/upload":
self._upload()
return
if path == "/api/training/mobile-packages":
self.close_connection = True
lengths = self.headers.get_all("Content-Length", [])
if (
self.headers.get("Transfer-Encoding")
or self.headers.get("Content-Encoding")
or self.headers.get("Content-Type") != "application/zip"
or len(lengths) != 1
or not re.fullmatch(r"[0-9]{1,10}", lengths[0])
):
raise ApiError(
HTTPStatus.BAD_REQUEST, "场景上传需要 application/zip 和唯一 Content-Length"
)
length = int(lengths[0])
if not 0 < length <= MAX_UPLOAD:
raise ApiError(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "场景上传上限 128 MiB")
self.connection.settimeout(60)
try:
result = self.manager.mobile_packages.receive(self.rfile, length)
except ValueError as error:
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
self._json(HTTPStatus.CREATED, result)
return
if path == "/api/training/jobs":
self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload()))
return
@@ -859,6 +1113,11 @@ def parse_args() -> argparse.Namespace:
parser.add_argument(
"--trainer-python", default=sys.executable, help="已安装 mjlab/torch 的 Python 解释器"
)
parser.add_argument(
"--mobile-python",
default=None,
help="可选独立移动操作 Python(MuJoCo 3.11.0 + SB3),避免更改 Go2 环境",
)
parser.add_argument(
"--tuning-data-root",
type=Path,
@@ -902,6 +1161,7 @@ def main() -> None:
args.trainer_root,
args.trainer_python,
tuple(args.tasks or DEFAULT_TASKS),
mobile_python=args.mobile_python,
lease=lease,
sources=sources,
)
+5 -1
View File
@@ -7,6 +7,8 @@ import math
import random
from typing import Any
from mobile_manipulator.config import MOBILE_TASKS, mobile_metadata
FLAT_TASK = "Unitree-Go2-Flat"
ROUGH_TASK = "Unitree-Go2-Rough"
OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance"
@@ -460,7 +462,9 @@ def task_metadata(tasks: tuple[str, ...]) -> list[dict]:
OBSTACLE_TASK: "前视射线避障导航",
}
return [
{
mobile_metadata(task)
if task in MOBILE_TASKS
else {
"id": task,
"name": names.get(task, task),
"browserCompatible": task in (FLAT_TASK, OBSTACLE_TASK),
@@ -0,0 +1,75 @@
"""Deterministic native-math oracle, consumed by Vitest. Run from repository root."""
import json
from pathlib import Path
import numpy as np
from training_server.mobile_manipulator.kernel import ROBOTS, TASK, TaskKernel, decode_action
from training_server.mobile_manipulator.motion import SafeActionController
def generate():
rng = np.random.default_rng(2026)
cases = []
for config in ROBOTS:
kernel = TaskKernel(config)
for i in range(20):
s = rng.uniform(-2, 2, TASK["stateSize"])
for start in [3, 33, 40, 50]:
q = rng.normal(size=4)
s[start : start + 4] = q / np.linalg.norm(q)
s[29] = rng.uniform(0, 1)
action = rng.uniform(-1.5, 1.5, TASK["actionSize"]).astype(np.float32)
kernel.reset()
kernel.has_lifted = bool(i % 2)
obs, reward, terminated, truncated, info = kernel.evaluate(s)
cases.append(
{
"robotId": config["id"],
"state": s.tolist(),
"action": action.tolist(),
"lifted": bool(i % 2),
"observation": obs.tolist(),
"control": decode_action(config, action).tolist(),
"reward": reward,
"terminated": terminated,
"truncated": truncated,
"info": info,
}
)
target = Path(__file__).resolve().parents[2] / "contracts/fixtures/mobile-golden.json"
target.write_text(json.dumps(cases, indent=2) + "\n")
motion_cases = []
for config in ROBOTS:
for stage in ("navigate", "reach", "pick-place"):
s = np.zeros(TASK["stateSize"])
s[[3, 33, 40, 50]] = 1
s[37:40] = [*TASK["navigationOffset"][:2], TASK["objectStart"][2]]
s[29] = 1
for i, joint in enumerate(config["armJoints"]):
s[13 + i] = joint["neutral"]
motion = SafeActionController(config)
motion.reset(s)
frames = []
for _ in range(12):
action = rng.uniform(-2, 2, TASK["actionSize"]).astype(np.float32)
motion.apply(action, s, stage, False)
frames.append(
dict(
action=action.tolist(),
control=motion.control.tolist(),
applied=motion.applied.tolist(),
targets=motion.targets.tolist(),
)
)
motion_cases.append(
dict(robotId=config["id"], stage=stage, state=s.tolist(), frames=frames)
)
target.with_name("mobile-motion-v2-golden.json").write_text(
json.dumps(motion_cases, indent=2) + "\n"
)
if __name__ == "__main__":
generate()
@@ -0,0 +1,224 @@
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()
@@ -0,0 +1,273 @@
"""One-click API/runner tests. No CUDA, mjlab or large robot assets required."""
import hashlib
import io
import json
import sys
import tempfile
import threading
import time
import unittest
import zipfile
from http.server import ThreadingHTTPServer
from pathlib import Path
from unittest.mock import patch
from urllib.error import HTTPError
from urllib.request import Request, urlopen
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from mobile_manipulator.config import MOBILE_TASKS, ROBOTS, TASK
from server import DEFAULT_TASKS, ApiError, TrainingJob, TrainingManager, TrainingRequestHandler
TASK_ID = next(iter(MOBILE_TASKS))
def archive(robot_id="lekiwi-v1", extra=None):
files = {
"robot.json": json.dumps(ROBOTS[robot_id], separators=(",", ":")),
"task.json": json.dumps(TASK),
"environment.json": json.dumps(
dict(scene="scene.xml", robotId=robot_id, taskId=TASK["id"], mujoco="3.11.0")
),
"scene.xml": "<mujoco><worldbody/></mujoco>",
}
files.update(extra or {})
stream = io.BytesIO()
with zipfile.ZipFile(stream, "w") as z:
for name, value in files.items():
z.writestr(zipfile.ZipInfo(name, date_time=(1980, 1, 1, 0, 0, 0)), value)
return stream.getvalue()
class MobileTrainingTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.manager = TrainingManager(
self.root, sys.executable, DEFAULT_TASKS, check_environment=False
)
raw = archive()
self.package = self.manager.mobile_packages.receive(io.BytesIO(raw), len(raw))
def tearDown(self):
self.manager.shutdown()
self.temp.cleanup()
def payload(self, **changes):
result = dict(
taskId=TASK_ID,
numEnvs=2,
maxIterations=2,
seed=123,
runName="mobile",
device="cpu",
gpuIds=[],
mobilePackageId=self.package["id"],
mobileParams=dict(rolloutSteps=8, goalPosition=[0.5, 0.4, 0.05]),
)
result.update(changes)
return result
def test_registry_readiness_is_per_family(self):
health = self.manager.health()
self.assertTrue(health["ready"])
mobile = [m for m in health["taskMetadata"] if m.get("family") == "mobile-manipulator"]
self.assertEqual(len(mobile), 2)
self.assertTrue(all(m["ready"] and not m["terrainPresets"] for m in mobile))
self.assertFalse(health["taskMetadata"][0]["ready"])
def test_validation_variant_and_parameter_bounds(self):
config = self.manager.parse_config(self.payload())
self.assertEqual(config.seed, 123)
self.assertEqual(config.mobile_params["rolloutSteps"], 8)
self.assertEqual(config.mobile_params["goalPosition"], [0.5, 0.4, 0.05])
self.assertEqual(config.deployment["sceneSha256"], self.package["sceneSha256"])
for fields in [
dict(taskId="MobileManipulator-LeKiwi-Bundle"),
dict(numEnvs=65),
dict(mobilePackageId="../../etc"),
dict(terrainPreset="plane"),
dict(mobileParams={"rolloutSteps": True}),
dict(mobileParams={"rolloutSteps": 7}),
dict(mobileParams={"stage": "fly"}),
dict(mobileParams={"stage": "reach"}),
dict(mobileParams={"sourceJobId": "../../untrusted.zip"}),
dict(mobileParams={"positionJitter": float("nan")}),
dict(mobileParams={"evaluationEpisodes": True}),
dict(mobileParams={"goalPosition": [0, float("nan"), 1]}),
dict(mobileParams={"goalPosition": [0, 0, -1]}),
dict(device="gpu", gpuIds=[0, 1]),
dict(pretrainedSourceId="abc"),
]:
with self.subTest(fields=fields), self.assertRaises(ApiError):
self.manager.parse_config(self.payload(**fields))
bundle = archive("lekiwi-bundle")
uploaded = self.manager.mobile_packages.receive(io.BytesIO(bundle), len(bundle))
self.assertEqual(
self.manager.parse_config(
self.payload(
taskId="MobileManipulator-LeKiwi-Bundle", mobilePackageId=uploaded["id"]
)
).deployment["robotId"],
"lekiwi-bundle",
)
def test_runner_arguments_use_server_owned_paths_and_mobile_interpreter(self):
self.manager.mobile_python = "/isolated/mobile/python"
args = self.manager.command_for(
self.manager.parse_config(self.payload(device="gpu", gpuIds=[2]))
)
self.assertEqual(args[0], self.manager.mobile_python)
self.assertIn("training_server.mobile_manipulator.train", args)
self.assertIn("cuda:2", args)
self.assertEqual(args[args.index("--seed") + 1], "123")
self.assertEqual(args[args.index("--num-envs") + 1], "2")
self.assertTrue(Path(args[args.index("--package") + 1]).is_relative_to(self.root))
def test_upload_rejects_traversal_xml_external_paths_and_wrong_contract(self):
for extra in [
{"../escape": "x"},
{"scene.xml": '<mujoco><include file="a.xml"/></mujoco>'},
{"scene.xml": '<mujoco><asset><mesh file="/etc/passwd"/></asset></mujoco>'},
{"scene.xml": '<mujoco><extension><plugin plugin="bad"/></extension></mujoco>'},
{"robot.json": "{}"},
{"robot.json": "[]"},
{"environment.json": "[]"},
{"task.json": "{}"},
{"./scene.xml": "<mujoco/>"},
{"scene.xml": '<mujoco><asset><texture fileleft="/etc/passwd"/></asset></mujoco>'},
]:
data = archive(extra=extra)
with self.subTest(extra=extra), self.assertRaises(ValueError):
self.manager.mobile_packages.receive(io.BytesIO(data), len(data))
with self.assertRaises(ValueError):
self.manager.mobile_packages.receive(io.BytesIO(b"bad"), 3)
with patch("mobile_manipulator.packages.MAX_EXPANDED", 1), self.assertRaises(ValueError):
data = archive(extra={"extra.txt": "other"})
self.manager.mobile_packages.receive(io.BytesIO(data), len(data))
def test_staged_resume_requires_matching_successful_evaluated_job(self):
source = TrainingJob(
id="d" * 32, config=self.manager.parse_config(self.payload()), state="succeeded"
)
source.artifact = self.root / "policy.onnx"
source.artifact.write_bytes(b"onnx")
source.artifact.with_suffix(".ppo.zip").write_bytes(b"trusted-checkpoint")
self.manager.jobs[source.id] = source
params = source.config.mobile_params | {"stage": "reach", "sourceJobId": source.id}
with self.assertRaisesRegex(ApiError, "尚未达标"):
self.manager.parse_config(self.payload(mobileParams=params))
source.config.deployment["evaluation"] = {
"episodes": 10,
"successRate": 0.8,
"safetyStops": 0,
}
config = self.manager.parse_config(self.payload(mobileParams=params))
command = self.manager.command_for(config)
self.assertEqual(
command[command.index("--resume") + 1], str(source.artifact.with_suffix(".ppo.zip"))
)
with self.assertRaisesRegex(ApiError, "依次推进"):
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "pick-place"}))
source.config.deployment["evaluation"]["safetyStops"] = 1
with self.assertRaisesRegex(ApiError, "尚未达标"):
self.manager.parse_config(self.payload(mobileParams=params))
# Poor quality still permits same-stage continuation, never a stage promotion.
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "navigate"}))
changed = archive(extra={"mesh.txt": "different bytes"})
other_package = self.manager.mobile_packages.receive(io.BytesIO(changed), len(changed))
with self.assertRaisesRegex(ApiError, "资产快照不匹配"):
self.manager.parse_config(
self.payload(
mobilePackageId=other_package["id"], mobileParams=params | {"stage": "navigate"}
)
)
source.config.deployment["sceneSha256"] = "other"
with self.assertRaisesRegex(ApiError, "不匹配"):
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "navigate"}))
def fake_command(self, config, _path=None):
# Real subprocess and lifecycle, deterministic stand-in only for expensive PPO/export.
script = self.root / "fake.py"
metadata = config.deployment | {"modelSha256": hashlib.sha256(b"onnx").hexdigest()}
script.write_text(
"import pathlib,sys,json,time\n"
"out=pathlib.Path(sys.argv[sys.argv.index('--output')+1])\n"
"print('Learning iteration 1 / 2',flush=True)\n"
"print('Mean surrogate loss: -0.25',flush=True)\n"
"time.sleep(.05)\n"
"out.write_bytes(b'onnx')\n"
f"out.with_name('deployment.json').write_text({json.dumps(json.dumps(metadata))})\n"
)
return [sys.executable, "-u", str(script)]
def test_http_create_poll_download_and_failures(self):
manager = self.manager
class Handler(TrainingRequestHandler):
access_token = "test-token"
def log_message(self, *_args):
pass
Handler.manager = manager
http = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=http.serve_forever, daemon=True)
thread.start()
base = f"http://127.0.0.1:{http.server_port}"
def request(path, body=None, content_type="application/json", token="test-token"):
req = Request(
base + path,
data=body,
headers={"Authorization": f"Bearer {token}", "Content-Type": content_type},
)
return urlopen(req, timeout=5)
try:
with self.assertRaises(HTTPError) as error:
request(
"/api/training/mobile-packages", archive(), "application/zip", token="wrong"
)
self.assertEqual(error.exception.code, 401)
error.exception.close()
with request("/api/training/mobile-packages", archive(), "application/zip") as response:
self.assertEqual(json.load(response)["id"], self.package["id"])
with patch.object(manager, "command_for", side_effect=self.fake_command):
with request("/api/training/jobs", json.dumps(self.payload()).encode()) as response:
self.assertEqual(response.status, 202)
job = json.load(response)
for _ in range(100):
with request("/api/training/jobs/" + job["id"]) as response:
job = json.load(response)
if job["state"] not in ("queued", "running"):
break
time.sleep(0.02)
self.assertEqual(job["state"], "succeeded", job)
self.assertEqual(job["progress"], 1)
self.assertIn("Mean surrogate loss: -0.25", job["logs"])
for filename in ("policy.onnx", "deployment.json"):
with request(f"/api/training/jobs/{job['id']}/artifacts/{filename}") as response:
self.assertTrue(response.read())
self.assertIsNone(manager.lease.public())
failed = TrainingJob(id="f" * 32, config=manager.parse_config(self.payload()))
with patch.object(
manager, "command_for", return_value=[sys.executable, "-c", "raise SystemExit(7)"]
):
manager._run(failed)
self.assertEqual(failed.state, "failed")
self.assertIsNone(failed.artifact)
finally:
http.shutdown()
http.server_close()
thread.join()
def test_cancel_before_launch_and_progress_parser(self):
job = TrainingJob(id="c" * 32, config=self.manager.parse_config(self.payload()))
self.manager._update_from_log(job, "\x1b[32mLearning iteration 1 / 2\x1b[0m")
self.assertEqual(job.public()["progress"], 0.5)
job.cancel_requested = True
with patch("server.subprocess.Popen") as popen:
self.manager._run(job)
popen.assert_not_called()
self.assertEqual(job.state, "cancelled")