From ecfcd3c869aec2b8b94c10e279e455e3c0afa955 Mon Sep 17 00:00:00 2001 From: cen617-code <1057290604@qq.com> Date: Thu, 27 Aug 2026 14:29:34 +0800 Subject: [PATCH] =?UTF-8?q?feat(web-platform):=20release=20V0.5.1=20?= =?UTF-8?q?=E5=89=8D=E7=AB=AF=20RL=20=E6=8E=A5=E5=85=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- wasm/package.json | 4 +- wasm/training_server/README.md | 43 ++ wasm/training_server/server.py | 492 ++++++++++++++++++ wasm/training_server/tests/test_server.py | 80 +++ wasm/web_platform/README.md | 17 +- .../components/LocalTrainingPanel.test.tsx | 25 + .../src/app/components/LocalTrainingPanel.tsx | 77 +++ .../src/app/components/SidebarPanel.tsx | 3 +- .../src/training/LocalTrainingClient.test.ts | 23 + .../src/training/LocalTrainingClient.ts | 36 ++ wasm/web_platform/src/training/types.ts | 40 ++ wasm/web_platform/src/viewer/MuJoCoViewer.ts | 31 +- .../src/viewer/interactionMath.test.ts | 27 + .../src/viewer/interactionMath.ts | 29 ++ 14 files changed, 911 insertions(+), 16 deletions(-) create mode 100644 wasm/training_server/README.md create mode 100644 wasm/training_server/server.py create mode 100644 wasm/training_server/tests/test_server.py create mode 100644 wasm/web_platform/src/app/components/LocalTrainingPanel.test.tsx create mode 100644 wasm/web_platform/src/app/components/LocalTrainingPanel.tsx create mode 100644 wasm/web_platform/src/training/LocalTrainingClient.test.ts create mode 100644 wasm/web_platform/src/training/LocalTrainingClient.ts create mode 100644 wasm/web_platform/src/training/types.ts create mode 100644 wasm/web_platform/src/viewer/interactionMath.test.ts create mode 100644 wasm/web_platform/src/viewer/interactionMath.ts diff --git a/wasm/package.json b/wasm/package.json index f7ecee29..e76270c0 100644 --- a/wasm/package.json +++ b/wasm/package.json @@ -23,7 +23,9 @@ "typecheck:platform": "tsc -p web_platform/tsconfig.json --noEmit", "lint:platform": "eslint web_platform/src web_platform/e2e", "test:platform": "vitest run --config web_platform/vite.config.ts", - "test:e2e:platform": "playwright test -c web_platform/playwright.config.ts" + "test:e2e:platform": "playwright test -c web_platform/playwright.config.ts", + "training-server": "python3 training_server/server.py", + "test:training-server": "python3 -m unittest discover -s training_server/tests" }, "author": "Google DeepMind", "license": "Apache-2.0", diff --git a/wasm/training_server/README.md b/wasm/training_server/README.md new file mode 100644 index 00000000..93e78dff --- /dev/null +++ b/wasm/training_server/README.md @@ -0,0 +1,43 @@ +# 本地强化学习训练服务 + +该服务把 Web 平台发出的受限训练请求转换为本机 `unitree_rl_mjlab/scripts/train.py` 子进程,并提供状态轮询、停止任务和 `policy.onnx` 下载接口。服务只绑定 `127.0.0.1`,不执行前端传入的任意命令。 + +## 启动 + +必须使用已经安装 `mjlab`、PyTorch 和训练工程依赖的 Python 解释器: + +```bash +npm run training-server --prefix wasm -- \ + --trainer-root /path/to/unitree_rl_mjlab \ + --trainer-python /path/to/training-env/bin/python +``` +如: +```bash +npm run training-server --prefix wasm -- --trainer-root /home/cen/Embodied_Workspace/unitree_rl_mjlab --trainer-python /home/cen/miniconda3/envs/unitree_rl_mjlab/bin/python +``` + +也可用 `UNITREE_RL_MJLAB_ROOT` 指定工程目录。默认端口是 `8765`。如果前端不是从 `localhost` 或 `127.0.0.1` 提供,可显式添加来源: + +```bash +python wasm/training_server/server.py \ + --trainer-root /path/to/unitree_rl_mjlab \ + --allow-origin http://192.168.1.10:5173 +``` + +服务一次只运行一个训练任务。停止服务或在界面点击“停止训练”会向整个训练进程组发送终止信号。训练请求的 W&B 模式默认为 `offline`,保留本地指标但不登录;也可以在界面选择完全禁用或在线模式。 + +## 接口 + +- `GET /api/training/health`:运行环境、允许的任务和活动任务; +- `POST /api/training/jobs`:发起训练; +- `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志; +- `DELETE /api/training/jobs/{id}`:停止训练; +- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略。 + +任务保存在服务内存中,服务重启后历史任务状态会丢失,但 mjlab 日志和 checkpoint 仍保留在训练工程中。 + +## 测试 + +```bash +npm run test:training-server --prefix wasm +``` diff --git a/wasm/training_server/server.py b/wasm/training_server/server.py new file mode 100644 index 00000000..161612c8 --- /dev/null +++ b/wasm/training_server/server.py @@ -0,0 +1,492 @@ +#!/usr/bin/env python3 +"""浏览器 MuJoCo 平台的本地 mjlab 训练桥接服务(仅绑定 loopback)。""" + +from __future__ import annotations + +import argparse +import json +import os +import re +import shutil +import signal +import subprocess +import sys +import threading +import time +import uuid +from collections import deque +from dataclasses import dataclass, field +from datetime import datetime, timezone +from http import HTTPStatus +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Any +from urllib.parse import unquote, urlsplit + +VERSION = "0.1.0" +# 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。 +DEFAULT_TASKS = ("Unitree-Go2-Flat",) +ACTIVE_STATES = {"queued", "running"} +ANSI_ESCAPE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]") +ITERATION_PATTERNS = ( + re.compile(r"(?:learning\s+)?iteration\D+(\d+)\s*/\s*(\d+)", re.I), + re.compile(r"(?:learning\s+)?iteration\D+(\d+)", re.I), +) +RUN_NAME = re.compile(r"^[A-Za-z0-9_.-]{1,64}$") +LOCAL_ORIGIN = re.compile(r"^https?://(?:localhost|127\.0\.0\.1)(?::\d+)?$") + + +def now_iso() -> str: + return datetime.now(timezone.utc).isoformat() + + +class ApiError(Exception): + def __init__(self, status: int, message: str): + super().__init__(message) + self.status = status + + +@dataclass +class TrainingConfig: + task_id: str + num_envs: int + max_iterations: int + seed: int + run_name: str + device: str + gpu_ids: list[int] + wandb_mode: str + + +@dataclass +class TrainingJob: + id: str + config: TrainingConfig + state: str = "queued" + created_at: str = field(default_factory=now_iso) + started_at: str | None = None + ended_at: str | None = None + iteration: int = 0 + message: str = "等待本地训练进程启动" + logs: deque[str] = field(default_factory=lambda: deque(maxlen=200)) + artifact: Path | None = None + process: subprocess.Popen[str] | None = None + cancel_requested: bool = False + + def public(self) -> dict[str, Any]: + progress = min(1.0, max(0.0, self.iteration / self.config.max_iterations)) + if self.state == "succeeded": + progress = 1.0 + return { + "id": self.id, + "state": self.state, + "taskId": self.config.task_id, + "createdAt": self.created_at, + "startedAt": self.started_at, + "endedAt": self.ended_at, + "iteration": self.iteration, + "maxIterations": self.config.max_iterations, + "progress": progress, + "message": self.message, + "logs": list(self.logs), + "artifactReady": self.artifact is not None and self.artifact.is_file(), + "artifactName": self.artifact.name if self.artifact else None, + } + + +class TrainingManager: + def __init__(self, trainer_root: Path, python: str, tasks: tuple[str, ...], check_environment: bool = True): + self.trainer_root = trainer_root.expanduser().resolve() + self.python = str(Path(python).expanduser()) if os.sep in python else python + self.tasks = tasks + self.jobs: dict[str, TrainingJob] = {} + self.lock = threading.RLock() + self.check_environment = check_environment + self._environment_error: str | None | bool = False + + def readiness_error(self) -> str | None: + if not self.trainer_root.is_dir(): + return f"训练工程目录不存在:{self.trainer_root}" + if not (self.trainer_root / "scripts" / "train.py").is_file(): + return f"训练入口不存在:{self.trainer_root / 'scripts/train.py'}" + executable = Path(self.python) + if not executable.is_file() and shutil.which(self.python) is None: + return f"Python 解释器不存在:{self.python}" + if self.check_environment and self._environment_error is False: + probe = "import importlib.util,sys; missing=[m for m in ('mjlab','torch','tyro') if importlib.util.find_spec(m) is None]; print(','.join(missing)); sys.exit(bool(missing))" + try: + result = subprocess.run([self.python, "-c", probe], cwd=self.trainer_root, capture_output=True, text=True, timeout=15, check=False) + missing = result.stdout.strip() + self._environment_error = f"训练 Python 缺少依赖:{missing}" if result.returncode else None + except (OSError, subprocess.TimeoutExpired) as error: + self._environment_error = f"无法检查训练 Python 环境:{error}" + return self._environment_error if isinstance(self._environment_error, str) else None + + def active_job_id(self) -> str | None: + with self.lock: + 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() + return { + "version": VERSION, + "ready": error is None, + "trainerRoot": str(self.trainer_root), + "python": self.python, + "tasks": list(self.tasks), + "activeJobId": self.active_job_id(), + "error": error, + } + + def parse_config(self, payload: Any) -> TrainingConfig: + if not isinstance(payload, dict): + raise ApiError(HTTPStatus.BAD_REQUEST, "请求体必须是 JSON 对象") + task_id = payload.get("taskId") + if task_id not in self.tasks: + raise ApiError(HTTPStatus.BAD_REQUEST, f"不允许的训练任务:{task_id}") + + def integer(name: str, minimum: int, maximum: int) -> int: + value = payload.get(name) + if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum: + raise ApiError(HTTPStatus.BAD_REQUEST, f"{name} 必须在 {minimum}–{maximum} 之间") + return value + + run_name = payload.get("runName", "web") + if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name): + raise ApiError(HTTPStatus.BAD_REQUEST, "runName 只能包含字母、数字、点、下划线和连字符,最长 64 字符") + device = payload.get("device") + if device not in ("cpu", "gpu"): + raise ApiError(HTTPStatus.BAD_REQUEST, "device 必须是 cpu 或 gpu") + raw_gpu_ids = payload.get("gpuIds", []) + if not isinstance(raw_gpu_ids, list) or any(isinstance(value, bool) or not isinstance(value, int) or value < 0 or value > 255 for value in raw_gpu_ids): + raise ApiError(HTTPStatus.BAD_REQUEST, "gpuIds 必须是非负整数数组") + if device == "gpu" and not raw_gpu_ids: + raise ApiError(HTTPStatus.BAD_REQUEST, "GPU 训练至少需要一个 GPU 编号") + wandb_mode = payload.get("wandbMode", "offline") + if wandb_mode not in ("offline", "disabled", "online"): + raise ApiError(HTTPStatus.BAD_REQUEST, "wandbMode 必须是 offline、disabled 或 online") + return TrainingConfig( + task_id=task_id, + num_envs=integer("numEnvs", 1, 16384), + max_iterations=integer("maxIterations", 1, 1_000_000), + seed=integer("seed", 0, 2_147_483_647), + run_name=run_name, + device=device, + gpu_ids=raw_gpu_ids, + wandb_mode=wandb_mode, + ) + + def start(self, payload: Any) -> dict[str, Any]: + error = self.readiness_error() + 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, "已有训练任务正在运行,请等待完成或先停止任务") + job = TrainingJob(id=uuid.uuid4().hex, config=config) + self.jobs[job.id] = job + threading.Thread(target=self._run, args=(job,), name=f"training-{job.id[:8]}", daemon=True).start() + return job.public() + + def get(self, job_id: str) -> dict[str, Any]: + with self.lock: + job = self.jobs.get(job_id) + if not job: + raise ApiError(HTTPStatus.NOT_FOUND, "训练任务不存在或服务已重启") + return job.public() + + def artifact(self, job_id: str) -> Path: + with self.lock: + job = self.jobs.get(job_id) + if not job: + raise ApiError(HTTPStatus.NOT_FOUND, "训练任务不存在或服务已重启") + if not job.artifact or not job.artifact.is_file(): + raise ApiError(HTTPStatus.NOT_FOUND, "该训练任务尚未生成 policy.onnx") + return job.artifact + + def cancel(self, job_id: str) -> dict[str, Any]: + with self.lock: + job = self.jobs.get(job_id) + if not job: + raise ApiError(HTTPStatus.NOT_FOUND, "训练任务不存在或服务已重启") + if job.state not in ACTIVE_STATES: + return job.public() + job.cancel_requested = True + job.message = "正在停止训练进程" + process = job.process + if process and process.poll() is None: + try: + os.killpg(process.pid, signal.SIGTERM) + except ProcessLookupError: + pass + threading.Thread(target=self._kill_later, args=(process,), daemon=True).start() + return self.get(job_id) + + @staticmethod + def _kill_later(process: subprocess.Popen[str]) -> None: + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + + def command_for(self, config: TrainingConfig) -> list[str]: + command = [ + self.python, "-u", "scripts/train.py", config.task_id, + f"--env.scene.num-envs={config.num_envs}", + f"--agent.max-iterations={config.max_iterations}", + f"--agent.seed={config.seed}", + f"--agent.run-name={config.run_name}", + ] + if config.device == "cpu": + command.extend(("--gpu-ids", "None")) + else: + # mjlab.TYRO_FLAGS 对 Union[list[int], Literal["all"], None] 使用 JSON 风格 + # list token;传成多个独立参数会被解析为错误的 Union 分支。 + command.extend(("--gpu-ids", json.dumps(config.gpu_ids, separators=(",", ":")))) + return command + + def _update_from_log(self, job: TrainingJob, raw_line: str) -> None: + line = ANSI_ESCAPE.sub("", raw_line).strip() + if not line: + return + with self.lock: + job.logs.append(line[-4000:]) + for pattern in ITERATION_PATTERNS: + match = pattern.search(line) + if match: + job.iteration = min(job.config.max_iterations, max(job.iteration, int(match.group(1)))) + break + job.message = line[-240:] + + def _artifact_snapshot(self) -> dict[Path, int]: + root = self.trainer_root / "logs" / "rsl_rl" + if not root.is_dir(): + return {} + return {path: path.stat().st_mtime_ns for path in root.glob("**/policy.onnx") if path.is_file()} + + def _find_artifact(self, before: dict[Path, int]) -> Path | None: + root = self.trainer_root / "logs" / "rsl_rl" + if not root.is_dir(): + return None + changed = [path for path in root.glob("**/policy.onnx") if path.is_file() and before.get(path) != path.stat().st_mtime_ns] + return max(changed, key=lambda path: path.stat().st_mtime_ns) if changed else None + + def _run(self, job: TrainingJob) -> None: + before = self._artifact_snapshot() + command = self.command_for(job.config) + with self.lock: + if job.cancel_requested: + job.state, job.ended_at, job.message = "cancelled", now_iso(), "训练已取消" + return + job.state, job.started_at, job.message = "running", now_iso(), "本地训练进程已启动" + try: + environment = os.environ.copy() + # 默认离线记录,保留本地 W&B 指标但不要求 API Key;只有前端明确选择 + # online 时才允许 wandb 发起登录/联网。 + environment["WANDB_MODE"] = job.config.wandb_mode + environment["WANDB_SILENT"] = "true" + process = subprocess.Popen( + command, + cwd=self.trainer_root, + env=environment, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + encoding="utf-8", + errors="replace", + bufsize=1, + start_new_session=True, + ) + with self.lock: + job.process = process + assert process.stdout is not None + try: + for line in process.stdout: + self._update_from_log(job, line) + finally: + process.stdout.close() + return_code = process.wait() + artifact = self._find_artifact(before) + with self.lock: + job.process = None + job.ended_at = now_iso() + if job.cancel_requested: + job.state, job.message = "cancelled", "训练已由用户取消" + elif return_code != 0: + job.state, job.message = "failed", f"训练进程退出,返回码 {return_code}" + elif artifact is None: + job.state, job.message = "failed", "训练结束,但没有找到本次生成的 policy.onnx" + else: + job.state, job.artifact = "succeeded", artifact + job.iteration = job.config.max_iterations + job.message = f"训练完成:{artifact.relative_to(self.trainer_root)}" + except Exception as error: # 服务必须保留错误供前端诊断。 + with self.lock: + job.process = None + job.ended_at = now_iso() + job.state = "cancelled" if job.cancel_requested else "failed" + job.message = f"启动训练失败:{error}" + job.logs.append(job.message) + + +class TrainingRequestHandler(BaseHTTPRequestHandler): + manager: TrainingManager + allowed_origins: tuple[str, ...] = () + server_version = "MuJoCoLocalTraining/0.1" + + def log_message(self, format: str, *args: Any) -> None: + sys.stderr.write(f"[{self.log_date_time_string()}] {format % args}\n") + + def _origin_allowed(self) -> bool: + origin = self.headers.get("Origin") + return origin is None or bool(LOCAL_ORIGIN.fullmatch(origin)) or origin in self.allowed_origins + + def _cors(self) -> None: + origin = self.headers.get("Origin") + if origin and self._origin_allowed(): + self.send_header("Access-Control-Allow-Origin", origin) + self.send_header("Vary", "Origin") + + def _json(self, status: int, payload: Any) -> None: + body = json.dumps(payload, ensure_ascii=False).encode("utf-8") + self.send_response(status) + self._cors() + self.send_header("Content-Type", "application/json; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.send_header("Cache-Control", "no-store") + self.end_headers() + self.wfile.write(body) + + def _error(self, error: Exception) -> None: + if isinstance(error, ApiError): + self._json(error.status, {"error": str(error)}) + else: + self._json(HTTPStatus.INTERNAL_SERVER_ERROR, {"error": f"本地训练服务内部错误:{error}"}) + + def _ensure_origin(self) -> None: + if not self._origin_allowed(): + raise ApiError(HTTPStatus.FORBIDDEN, "不允许的浏览器来源") + + def _payload(self) -> Any: + try: + length = int(self.headers.get("Content-Length", "0")) + except ValueError as error: + raise ApiError(HTTPStatus.BAD_REQUEST, "Content-Length 无效") from error + if length <= 0 or length > 32 * 1024: + raise ApiError(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "训练请求体不能为空且不能超过 32 KiB") + try: + return json.loads(self.rfile.read(length)) + except (UnicodeDecodeError, json.JSONDecodeError) as error: + raise ApiError(HTTPStatus.BAD_REQUEST, "训练请求不是有效 JSON") from error + + @staticmethod + def _route(path: str) -> tuple[str | None, bool]: + match = re.fullmatch(r"/api/training/jobs/([0-9a-f]{32})(/artifacts/policy\.onnx)?", path) + return (unquote(match.group(1)), bool(match.group(2))) if match else (None, False) + + def do_OPTIONS(self) -> None: + try: + self._ensure_origin() + self.send_response(HTTPStatus.NO_CONTENT) + self._cors() + self.send_header("Access-Control-Allow-Methods", "GET, POST, DELETE, OPTIONS") + self.send_header("Access-Control-Allow-Headers", "Content-Type") + self.send_header("Access-Control-Max-Age", "600") + self.end_headers() + except Exception as error: + self._error(error) + + def do_GET(self) -> None: + try: + self._ensure_origin() + path = urlsplit(self.path).path + if path == "/api/training/health": + self._json(HTTPStatus.OK, self.manager.health()) + return + job_id, artifact = self._route(path) + if not job_id: + raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") + if artifact: + file_path = self.manager.artifact(job_id) + size = file_path.stat().st_size + self.send_response(HTTPStatus.OK) + self._cors() + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Disposition", 'attachment; filename="policy.onnx"') + self.send_header("Content-Length", str(size)) + self.send_header("Cache-Control", "no-store") + self.end_headers() + with file_path.open("rb") as source: + shutil.copyfileobj(source, self.wfile) + else: + self._json(HTTPStatus.OK, self.manager.get(job_id)) + except Exception as error: + self._error(error) + + def do_POST(self) -> None: + try: + self._ensure_origin() + if urlsplit(self.path).path != "/api/training/jobs": + raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") + self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload())) + except Exception as error: + self._error(error) + + def do_DELETE(self) -> None: + try: + self._ensure_origin() + job_id, artifact = self._route(urlsplit(self.path).path) + if not job_id or artifact: + raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") + self._json(HTTPStatus.ACCEPTED, self.manager.cancel(job_id)) + except Exception as error: + self._error(error) + + +def default_trainer_root() -> Path: + configured = os.environ.get("UNITREE_RL_MJLAB_ROOT") + if configured: + return Path(configured) + repository = Path(__file__).resolve().parents[2] + return repository.parent.parent / "unitree_rl_mjlab" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="MuJoCo Web 平台本地强化学习训练服务") + parser.add_argument("--host", default="127.0.0.1", choices=("127.0.0.1", "localhost"), help="仅允许绑定本机回环地址") + parser.add_argument("--port", type=int, default=8765) + parser.add_argument("--trainer-root", type=Path, default=default_trainer_root(), help="unitree_rl_mjlab 工程目录") + parser.add_argument("--trainer-python", default=sys.executable, help="已安装 mjlab/torch 的 Python 解释器") + parser.add_argument("--task", action="append", dest="tasks", help="允许前端启动的任务 ID;可重复") + parser.add_argument("--allow-origin", action="append", default=[], help="额外允许的前端 Origin;可重复") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + manager = TrainingManager(args.trainer_root, args.trainer_python, tuple(args.tasks or DEFAULT_TASKS)) + TrainingRequestHandler.manager = manager + TrainingRequestHandler.allowed_origins = tuple(args.allow_origin) + server = ThreadingHTTPServer((args.host, args.port), TrainingRequestHandler) + print(f"本地训练服务:http://{args.host}:{args.port}") + print(f"训练工程:{manager.trainer_root}") + print(f"Python:{manager.python}") + if manager.readiness_error(): + print(f"警告:{manager.readiness_error()}", file=sys.stderr) + try: + server.serve_forever() + except KeyboardInterrupt: + print("\n正在停止本地训练服务…") + finally: + active = manager.active_job_id() + if active: + manager.cancel(active) + server.server_close() + + +if __name__ == "__main__": + main() diff --git a/wasm/training_server/tests/test_server.py b/wasm/training_server/tests/test_server.py new file mode 100644 index 00000000..ef2ad2c1 --- /dev/null +++ b/wasm/training_server/tests/test_server.py @@ -0,0 +1,80 @@ +import sys +import tempfile +import time +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from server import ApiError, TrainingManager # noqa: E402 + + +class TrainingManagerTest(unittest.TestCase): + def setUp(self): + self.temporary = tempfile.TemporaryDirectory() + self.root = Path(self.temporary.name) + (self.root / "scripts").mkdir() + (self.root / "scripts" / "train.py").write_text( + """import os, pathlib, time +print('WANDB_MODE=' + os.environ.get('WANDB_MODE', ''), flush=True) +print('Learning iteration 1 / 2', flush=True) +time.sleep(0.02) +print('Learning iteration 2 / 2', flush=True) +out=pathlib.Path('logs/rsl_rl/test/run/policy.onnx') +out.parent.mkdir(parents=True, exist_ok=True) +out.write_bytes(b'onnx') +""", + encoding="utf-8", + ) + self.manager = TrainingManager(self.root, sys.executable, ("Unitree-Go2-Flat",), check_environment=False) + + def tearDown(self): + self.temporary.cleanup() + + @staticmethod + def payload(**patch): + value = { + "taskId": "Unitree-Go2-Flat", + "numEnvs": 16, + "maxIterations": 2, + "seed": 42, + "runName": "browser-test", + "device": "cpu", + "gpuIds": [], + "wandbMode": "offline", + } + value.update(patch) + return value + + def test_validates_allowlist_and_limits(self): + with self.assertRaises(ApiError): + self.manager.parse_config(self.payload(taskId="shell injection")) + with self.assertRaises(ApiError): + self.manager.parse_config(self.payload(numEnvs=0)) + with self.assertRaises(ApiError): + self.manager.parse_config(self.payload(runName="bad name")) + with self.assertRaises(ApiError): + self.manager.parse_config(self.payload(wandbMode="login")) + + def test_builds_argument_array_without_shell(self): + config = self.manager.parse_config(self.payload(device="gpu", gpuIds=[0, 2])) + command = self.manager.command_for(config) + self.assertEqual(command[:4], [sys.executable, "-u", "scripts/train.py", "Unitree-Go2-Flat"]) + self.assertEqual(command[-2:], ["--gpu-ids", "[0,2]"]) + + def test_runs_job_and_exposes_new_onnx_artifact(self): + job = self.manager.start(self.payload()) + deadline = time.monotonic() + 5 + while time.monotonic() < deadline: + job = self.manager.get(job["id"]) + if job["state"] not in ("queued", "running"): + break + time.sleep(0.02) + self.assertEqual(job["state"], "succeeded") + self.assertEqual(job["iteration"], 2) + self.assertIn("WANDB_MODE=offline", job["logs"]) + self.assertTrue(job["artifactReady"]) + self.assertEqual(self.manager.artifact(job["id"]).read_bytes(), b"onnx") + + +if __name__ == "__main__": + unittest.main() diff --git a/wasm/web_platform/README.md b/wasm/web_platform/README.md index 51b2ff10..bf8b8639 100644 --- a/wasm/web_platform/README.md +++ b/wasm/web_platform/README.md @@ -14,6 +14,7 @@ - actuator 滑杆、hinge/slide 关节拖动、动态 body 外力拖拽 - 导入单文件 `.py` 控制器,通过本地 Pyodide 在 `mj_step` 前按仿真时间同步执行 - 导入 mjlab 导出的 `policy.onnx`,在浏览器本地执行 Go2-W 平衡/速度策略推理 +- 从图形界面向本机训练桥接服务发起 mjlab 强化学习训练、查看进度/日志、停止任务并导入训练生成的 ONNX - FPS、物理耗时和主线程步进预算提示 ## 开发 @@ -63,9 +64,23 @@ Python 控制器是可信的单文件脚本,必须同步定义 `step(ctx, stat 当前 Python 与 MuJoCo 都运行在主线程,以保证闭环调用严格位于 `mj_step` 前。仅运行可信脚本;死循环仍可能阻塞页面。Pyodide 及 Python 标准库由 npm 包随生产构建离线发布,不从 CDN 下载;暂不支持第三方 Python 包、`pip` 或多文件 import。 +## 本地强化学习训练 + +训练仍由本机 Python/mjlab 进程执行,但可以从右侧“控制 → 本地强化学习训练”直接发起和管理。先使用安装了 mjlab、PyTorch 及训练依赖的 Python 启动本地桥接服务: + +```bash +npm run training-server --prefix wasm -- \ + --trainer-root /path/to/unitree_rl_mjlab \ + --trainer-python /path/to/training-env/bin/python +``` + +界面默认连接 `http://127.0.0.1:8765`,可选择服务端允许的任务、并行环境数、训练迭代、随机种子、CPU/GPU、GPU 编号和实验记录方式。W&B 默认为本地离线模式,无需登录或 API Key;也可完全禁用,只有明确选择在线模式时才会联网登录。训练期间页面轮询迭代进度与最近日志,可以停止任务;训练成功后点击“导入策略”,生成的 `policy.onnx` 会进入现有 ONNX 加载流程。 + +桥接服务只监听本机回环地址、仅接受允许列表中的任务和经过范围校验的参数,不执行前端提供的 Shell 命令;一次只运行一个训练进程。当前任务使用 `unitree_rl_mjlab` 自带的机器人资产与环境配置,**不会自动把浏览器中临时编辑的 MJCF/URDF 作为训练环境**。自定义浏览器模型训练需要先在 mjlab 中注册对应 task。服务配置、接口和安全边界见 [`../training_server/README.md`](../training_server/README.md)。 + ## ONNX 强化学习策略 -平台只负责策略推理,训练仍在 Python/mjlab 中完成。当前内置任务兼容 `unitree_rl_mjlab` Go2 velocity 的部署观测顺序: +当前内置任务兼容 `unitree_rl_mjlab` Go2 velocity 的部署观测顺序: ```text base_ang_vel(3) + projected_gravity(3) + velocity_command(3) diff --git a/wasm/web_platform/src/app/components/LocalTrainingPanel.test.tsx b/wasm/web_platform/src/app/components/LocalTrainingPanel.test.tsx new file mode 100644 index 00000000..462a6529 --- /dev/null +++ b/wasm/web_platform/src/app/components/LocalTrainingPanel.test.tsx @@ -0,0 +1,25 @@ +import {fireEvent,render,screen,waitFor} from '@testing-library/react'; +import {beforeEach,describe,expect,it,vi} from 'vitest'; +import {LocalTrainingPanel} from './LocalTrainingPanel'; + +beforeEach(()=>{localStorage.clear();vi.unstubAllGlobals();}); + +describe('LocalTrainingPanel',()=>{ + it('连接本地服务并从图形界面发起训练请求',async()=>{ + const health={version:'0.1.0',ready:true,trainerRoot:'/opt/unitree_rl_mjlab',python:'/env/bin/python',tasks:['Unitree-Go2-Flat']}; + const job={id:'a'.repeat(32),state:'queued',taskId:'Unitree-Go2-Flat',createdAt:'2025-01-01T00:00:00Z',iteration:0,maxIterations:2000,progress:0,message:'等待启动',logs:[],artifactReady:false}; + const fetchMock=vi.fn() + .mockResolvedValueOnce(new Response(JSON.stringify(health),{status:200,headers:{'Content-Type':'application/json'}})) + .mockResolvedValueOnce(new Response(JSON.stringify(job),{status:202,headers:{'Content-Type':'application/json'}})); + vi.stubGlobal('fetch',fetchMock); + render(); + fireEvent.click(screen.getByRole('button',{name:'连接'})); + expect(await screen.findByText('/opt/unitree_rl_mjlab')).toBeInTheDocument(); + fireEvent.change(screen.getByLabelText('并行环境'),{target:{value:'32'}}); + fireEvent.click(screen.getByRole('button',{name:'发起本地训练'})); + await waitFor(()=>expect(fetchMock).toHaveBeenCalledTimes(2)); + const request=fetchMock.mock.calls[1][1] as RequestInit; + expect(JSON.parse(String(request.body))).toMatchObject({taskId:'Unitree-Go2-Flat',numEnvs:32,device:'gpu',gpuIds:[0],wandbMode:'offline'}); + expect(await screen.findByText('排队中')).toBeInTheDocument(); + }); +}); diff --git a/wasm/web_platform/src/app/components/LocalTrainingPanel.tsx b/wasm/web_platform/src/app/components/LocalTrainingPanel.tsx new file mode 100644 index 00000000..f61be61d --- /dev/null +++ b/wasm/web_platform/src/app/components/LocalTrainingPanel.tsx @@ -0,0 +1,77 @@ +import {useEffect,useState,type ReactNode} from 'react'; +import {Download,Link,Play,Server,Square} from 'lucide-react'; +import {Badge,Button,ProgressBar,PropertyRow,Select} from '../../components/ui'; +import {LocalTrainingClient} from '../../training/LocalTrainingClient'; +import type {TrainingDevice,TrainingJob,TrainingServerInfo,WandbMode} from '../../training/types'; + +const ENDPOINT_KEY='mujoco-local-training-endpoint',JOB_KEY='mujoco-local-training-job'; +const DEFAULT_ENDPOINT='http://127.0.0.1:8765'; +const ACTIVE_STATES=new Set(['queued','running']); +function stored(key:string,fallback=''):string{try{return localStorage.getItem(key)??fallback;}catch{return fallback;}} +function errorText(error:unknown):string{return error instanceof Error?error.message:String(error);} +function stateLabel(state:TrainingJob['state']):string{return {queued:'排队中',running:'训练中',succeeded:'已完成',failed:'失败',cancelled:'已取消'}[state];} + +export function LocalTrainingPanel({onPolicyReady}:{onPolicyReady(file:File):void}){ + const [endpoint,setEndpoint]=useState(()=>stored(ENDPOINT_KEY,DEFAULT_ENDPOINT)); + const [server,setServer]=useState(); + const [job,setJob]=useState(); + const [busy,setBusy]=useState(false),[error,setError]=useState(); + const [taskId,setTaskId]=useState('Unitree-Go2-Flat'),[numEnvs,setNumEnvs]=useState(4096),[maxIterations,setMaxIterations]=useState(2000),[seed,setSeed]=useState(42),[runName,setRunName]=useState('web'),[device,setDevice]=useState('gpu'),[gpuIds,setGpuIds]=useState('0'),[wandbMode,setWandbMode]=useState('offline'); + + const connect=async()=>{ + setBusy(true);setError(undefined); + try{ + const client=new LocalTrainingClient(endpoint),info=await client.health(); + setServer(info);try{localStorage.setItem(ENDPOINT_KEY,client.endpoint);}catch{/* 当前会话仍可连接 */} + if(info.tasks.length&&!info.tasks.includes(taskId))setTaskId(info.tasks[0]); + const remembered=info.activeJobId??stored(JOB_KEY); + if(remembered){try{setJob(await client.job(remembered));}catch{try{localStorage.removeItem(JOB_KEY);}catch{/* ignore */}}} + if(!info.ready)setError(info.error??'训练服务尚未就绪'); + }catch(value){setServer(undefined);setError(errorText(value));} + finally{setBusy(false);} + }; + + const jobId=job?.id,jobState=job?.state; + useEffect(()=>{ + if(!jobId||!jobState||!ACTIVE_STATES.has(jobState))return; + let disposed=false; + const refresh=async()=>{try{const next=await new LocalTrainingClient(endpoint).job(jobId);if(!disposed)setJob(next);}catch(value){if(!disposed)setError(errorText(value));}}; + const timer=window.setInterval(()=>void refresh(),1500);return()=>{disposed=true;window.clearInterval(timer);}; + },[endpoint,jobId,jobState]); + + const start=async()=>{ + setBusy(true);setError(undefined); + try{ + const ids=device==='gpu'?gpuIds.split(/[\s,]+/).filter(Boolean).map(Number):[]; + if(ids.some(id=>!Number.isInteger(id)||id<0))throw new Error('GPU 编号必须是非负整数'); + const next=await new LocalTrainingClient(endpoint).start({taskId,numEnvs,maxIterations,seed,runName,device,gpuIds:ids,wandbMode}); + setJob(next);try{localStorage.setItem(JOB_KEY,next.id);}catch{/* ignore */} + }catch(value){setError(errorText(value));}finally{setBusy(false);} + }; + const cancel=async()=>{if(!job)return;setBusy(true);setError(undefined);try{setJob(await new LocalTrainingClient(endpoint).cancel(job.id));}catch(value){setError(errorText(value));}finally{setBusy(false);}}; + const importResult=async()=>{if(!job)return;setBusy(true);setError(undefined);try{onPolicyReady(await new LocalTrainingClient(endpoint).downloadPolicy(job.id));}catch(value){setError(errorText(value));}finally{setBusy(false);}}; + const active=Boolean(job&&ACTIVE_STATES.has(job.state)); + + return
+ +
{server?.trainerRoot??'请先启动本地训练服务'}{server?.ready?'可用':'离线'}
+ {server?.ready&&!job&&
+ +
setRunName(event.target.value)}/>
+
setGpuIds(event.target.value)}/>
+ + +

训练使用本地 mjlab 任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。

+
} + {job&&
+
{job.taskId}{stateLabel(job.state)}
+
+ {job.logs.length>0&&
最近日志
{job.logs.slice(-40).join('\n')}
} +
{active?:<>}
+
} + {error&&

{error}

} +
; +} + +function Field({label,children}:{label:string;children:ReactNode}){return ;} +function NumberField({label,value,min,max,onChange}:{label:string;value:number;min:number;max:number;onChange(value:number):void}){return onChange(Number(event.target.value))}/>;} diff --git a/wasm/web_platform/src/app/components/SidebarPanel.tsx b/wasm/web_platform/src/app/components/SidebarPanel.tsx index 67959bc5..e1a488af 100644 --- a/wasm/web_platform/src/app/components/SidebarPanel.tsx +++ b/wasm/web_platform/src/app/components/SidebarPanel.tsx @@ -13,6 +13,7 @@ import {TreeSearchField} from './TreeSearchField'; import {ProjectBreadcrumb} from './ProjectBreadcrumb'; import {PythonControllerPanel} from './PythonControllerPanel'; import {RLPolicyPanel} from './RLPolicyPanel'; +import {LocalTrainingPanel} from './LocalTrainingPanel'; export function SidebarPanel({title,side,children,visible=true}:{title:string;side:'left'|'right';children:ReactNode;visible?:boolean}){return ;} @@ -33,7 +34,7 @@ export function ModelControlsSidebar(props:ModelControlsProps){const [tab,setTab const properties=<>{s.model.nbody} Body}>
{props.selectedFormat==='urdf'&&

MJCF 模式保留 visual mesh、添加物理地面,并将模型最低点对齐到 z=0。

} {props.selection?
}/>}/>value.toFixed(3)).join(', ')} action={}/>
:

在视口中单击物体

}
; - const controls=<>{s.rlPolicy.enabled?'推理':'停止'}:undefined}>{s.controller.enabled?'运行':'停止'}:undefined}>{s.actuators.length}}>{s.actuators.length?s.actuators.map(actuator=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):

模型没有驱动器

}
+ const controls=<>{s.rlPolicy.enabled?'推理':'停止'}:undefined}>{s.controller.enabled?'运行':'停止'}:undefined}>{s.actuators.length}}>{s.actuators.length?s.actuators.map(actuator=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):

模型没有驱动器

}
{s.joints.length}}>
{s.joints.map(joint=>{const scale=joint.type===3&&props.angleUnit==='deg'?180/Math.PI:1,unit=joint.type===3?(props.angleUnit==='deg'?'°':' rad'):joint.type===2?' m':'';return props.onJoint(joint.id,value/scale)}/>;})}

选择“外力施加”,在动态物体上按住拖动,松开即清零。

; return ,content:properties},{value:'controls',label:'控制',icon:,content:controls}]}/>; diff --git a/wasm/web_platform/src/training/LocalTrainingClient.test.ts b/wasm/web_platform/src/training/LocalTrainingClient.test.ts new file mode 100644 index 00000000..01f8bbaf --- /dev/null +++ b/wasm/web_platform/src/training/LocalTrainingClient.test.ts @@ -0,0 +1,23 @@ +import {afterEach,describe,expect,it,vi} from 'vitest'; +import {LocalTrainingClient} from './LocalTrainingClient'; + +afterEach(()=>vi.unstubAllGlobals()); + +describe('LocalTrainingClient',()=>{ + it('规范化服务地址并提交受类型约束的 JSON 请求',async()=>{ + const fetchMock=vi.fn().mockResolvedValue(new Response(JSON.stringify({id:'a'.repeat(32),state:'queued'}),{status:202,headers:{'Content-Type':'application/json'}})); + vi.stubGlobal('fetch',fetchMock); + const client=new LocalTrainingClient('http://127.0.0.1:8765/'); + await client.start({taskId:'Unitree-Go2-Flat',numEnvs:16,maxIterations:2,seed:42,runName:'test',device:'cpu',gpuIds:[],wandbMode:'offline'}); + expect(fetchMock).toHaveBeenCalledWith('http://127.0.0.1:8765/api/training/jobs',expect.objectContaining({method:'POST'})); + const options=fetchMock.mock.calls[0][1] as RequestInit; + expect(JSON.parse(String(options.body))).toMatchObject({taskId:'Unitree-Go2-Flat',numEnvs:16,device:'cpu'}); + }); + + it('显示服务端返回的中文错误',async()=>{ + vi.stubGlobal('fetch',vi.fn().mockResolvedValue(new Response(JSON.stringify({error:'已有训练任务正在运行'}),{status:409,headers:{'Content-Type':'application/json'}}))); + await expect(new LocalTrainingClient('http://localhost:8765').health()).rejects.toThrow('已有训练任务正在运行'); + }); + + it('拒绝非 HTTP 地址',()=>{expect(()=>new LocalTrainingClient('file:///tmp/socket')).toThrow('http 或 https');}); +}); diff --git a/wasm/web_platform/src/training/LocalTrainingClient.ts b/wasm/web_platform/src/training/LocalTrainingClient.ts new file mode 100644 index 00000000..dff2eaa5 --- /dev/null +++ b/wasm/web_platform/src/training/LocalTrainingClient.ts @@ -0,0 +1,36 @@ +import type {TrainingJob,TrainingRequest,TrainingServerInfo} from './types'; + +function normalizeEndpoint(value:string):string{ + const endpoint=value.trim().replace(/\/+$/,''); + let url:URL; + try{url=new URL(endpoint);}catch{throw new Error('训练服务地址无效');} + if(url.protocol!=='http:'&&url.protocol!=='https:')throw new Error('训练服务地址必须使用 http 或 https'); + return url.toString().replace(/\/$/,''); +} + +async function responseError(response:Response):Promise{ + try{const body=await response.json() as {error?:string};if(body.error)return new Error(body.error);}catch{/* 使用 HTTP 状态作为回退 */} + return new Error(`本地训练服务请求失败(HTTP ${response.status})`); +} + +export class LocalTrainingClient { + readonly endpoint:string; + constructor(endpoint:string){this.endpoint=normalizeEndpoint(endpoint);} + + private async json(path:string,init?:RequestInit):Promise{ + const response=await fetch(`${this.endpoint}${path}`,init); + if(!response.ok)throw await responseError(response); + return response.json() as Promise; + } + + health():Promise{return this.json('/api/training/health');} + start(request:TrainingRequest):Promise{return this.json('/api/training/jobs',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify(request)});} + job(id:string):Promise{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`);} + cancel(id:string):Promise{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`,{method:'DELETE'});} + async downloadPolicy(id:string):Promise{ + const response=await fetch(`${this.endpoint}/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`); + if(!response.ok)throw await responseError(response); + const blob=await response.blob(); + return new File([blob],`policy-${id.slice(0,8)}.onnx`,{type:'application/octet-stream'}); + } +} diff --git a/wasm/web_platform/src/training/types.ts b/wasm/web_platform/src/training/types.ts new file mode 100644 index 00000000..03195fdf --- /dev/null +++ b/wasm/web_platform/src/training/types.ts @@ -0,0 +1,40 @@ +export type TrainingJobState='queued'|'running'|'succeeded'|'failed'|'cancelled'; +export type TrainingDevice='cpu'|'gpu'; +export type WandbMode='offline'|'online'|'disabled'; + +export interface TrainingServerInfo { + version:string; + ready:boolean; + trainerRoot:string; + python:string; + tasks:string[]; + activeJobId?:string; + error?:string; +} + +export interface TrainingRequest { + taskId:string; + numEnvs:number; + maxIterations:number; + seed:number; + runName:string; + device:TrainingDevice; + gpuIds:number[]; + wandbMode:WandbMode; +} + +export interface TrainingJob { + id:string; + state:TrainingJobState; + taskId:string; + createdAt:string; + startedAt?:string; + endedAt?:string; + iteration:number; + maxIterations:number; + progress:number; + message:string; + logs:string[]; + artifactReady:boolean; + artifactName?:string; +} diff --git a/wasm/web_platform/src/viewer/MuJoCoViewer.ts b/wasm/web_platform/src/viewer/MuJoCoViewer.ts index 884ef4ba..10338c85 100644 --- a/wasm/web_platform/src/viewer/MuJoCoViewer.ts +++ b/wasm/web_platform/src/viewer/MuJoCoViewer.ts @@ -4,6 +4,7 @@ import type {MjvGeom, MjvOption, MjvCamera, MjvScene} from '@mujoco/mujoco'; import type {FrameResult, SimulationSession, SimulationSnapshot} from '../simulation/SimulationSession'; import {meshIdFromSceneDataId} from '../simulation/geometry'; import {OrientationGizmo} from './OrientationGizmo'; +import {cameraAlignedForce,closestRayAxisParameter,resolveHingeDragDelta,signedAngleAroundAxis} from './interactionMath'; export type InteractionMode = 'select' | 'joint' | 'force'; export type ViewerTheme='light'|'dark'; @@ -31,7 +32,8 @@ export class MuJoCoViewer { private mjScene: MjvScene | null = null; private frame = 0; private lastFpsAt=performance.now(); private fpsFrames=0; private lastSnapshotAt=0; private meshes: THREE.Mesh[]=[]; private geometries=new Map(); private textures=new Map(); - private raycaster=new THREE.Raycaster(); private pointer=new THREE.Vector2(); private selected:THREE.Mesh|null=null; private dragStart:THREE.Vector2|null=null; private dragAxis=new THREE.Vector2(1,0); private dragJointId=-1; private dragJointValue=0; private arrow:THREE.ArrowHelper|null=null;private highlightedJointId=-1;private highlightedBodyId=-1;private jointMarker:THREE.Mesh|null=null; + private readonly sensorRotation=new THREE.Matrix4();private readonly renderSize=new THREE.Vector2(); + private raycaster=new THREE.Raycaster(); private pointer=new THREE.Vector2(); private selected:THREE.Mesh|null=null; private dragStart:THREE.Vector2|null=null; private dragJointId=-1; private dragJointType=-1;private dragJointValue=0;private dragHitDistance=0;private dragSlideParameter=0;private dragJointPivot=new THREE.Vector3();private dragJointAxisWorld=new THREE.Vector3();private dragJointStartWorld=new THREE.Vector3();private dragJointStartPlaneVector=new THREE.Vector3(); private arrow:THREE.ArrowHelper|null=null;private highlightedJointId=-1;private highlightedBodyId=-1;private jointMarker:THREE.Mesh|null=null; private resizeObserver:ResizeObserver; private orientationGizmo:OrientationGizmo; private grid:THREE.GridHelper; @@ -43,7 +45,7 @@ export class MuJoCoViewer { this.camera.up.set(0,0,1); this.camera.position.set(3,-3,2); this.controls=new OrbitControls(this.camera,this.renderer.domElement); this.controls.enableDamping=true; this.scene.background=new THREE.Color(0x0b1220);this.hemisphere=new THREE.HemisphereLight(0xffffff,0x223344,1.3);this.scene.add(this.hemisphere); const light=new THREE.DirectionalLight(0xffffff,2); light.position.set(4,-3,7); light.castShadow=true; this.scene.add(light);this.grid=new THREE.GridHelper(20,40,0x3b82f6,0x253047).rotateX(Math.PI/2);this.scene.add(this.grid);this.orientationGizmo=new OrientationGizmo(host);this.orientationGizmo.update(this.camera); this.resizeObserver=new ResizeObserver(()=>this.resize()); this.resizeObserver.observe(host); this.resize(); - this.renderer.domElement.addEventListener('pointerdown',this.onPointerDown); this.renderer.domElement.addEventListener('pointermove',this.onPointerMove); window.addEventListener('pointerup',this.onPointerUp); + this.renderer.domElement.addEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.addEventListener('pointermove',this.onPointerMove);this.renderer.domElement.addEventListener('lostpointercapture',this.onPointerUp);window.addEventListener('pointerup',this.onPointerUp);window.addEventListener('pointercancel',this.onPointerUp);window.addEventListener('blur',this.onPointerUp); this.frame=requestAnimationFrame(this.animate); } @@ -62,28 +64,31 @@ export class MuJoCoViewer { private fitCamera(session:SimulationSession):void {const {extent,center}=session.geometryBounds();this.controls.target.set(center[0],center[1],center[2]);this.camera.position.set(center[0]+extent*1.5,center[1]-extent*1.5,center[2]+extent);this.camera.near=Math.max(.001,extent/1000);this.camera.far=Math.max(100,extent*100);this.camera.updateProjectionMatrix();this.controls.update();} private resize():void {const w=Math.max(1,this.host.clientWidth),h=Math.max(1,this.host.clientHeight); this.renderer.setSize(w,h,false); this.camera.aspect=w/h; this.camera.updateProjectionMatrix();} - private animate=(now:number):void=>{try {const result=this.session?.advance(now)??{steps:0,stepMs:0,overBudget:false};this.updateThemeTransition(now); this.controls.update();this.orientationGizmo.update(this.camera); if(this.session){this.updateMuJoCoScene();this.updateJointMarker();} this.renderer.setScissorTest(false);this.renderer.render(this.scene,this.camera);if(this.showSensorCamera&&this.updateSensorCamera())this.renderSensorCamera(); this.fpsFrames++; let fps=0;if(now-this.lastFpsAt>=500){fps=this.fpsFrames*1000/(now-this.lastFpsAt);this.fpsFrames=0;this.lastFpsAt=now;} const snapshot=this.session&&now-this.lastSnapshotAt>150?(this.lastSnapshotAt=now,this.session.snapshot()):undefined; this.callbacks.onFrame(result,fps,snapshot);}catch(error){this.callbacks.onError(error instanceof Error?error:new Error(String(error)));} this.frame=requestAnimationFrame(this.animate);}; + private animate=(now:number):void=>{try {const result=this.session?.advance(now)??{steps:0,stepMs:0,overBudget:false};this.updateThemeTransition(now); this.controls.update();this.orientationGizmo.update(this.camera); if(this.session){this.updateMuJoCoScene();this.updateJointMarker();} this.renderer.setScissorTest(false);this.renderer.render(this.scene,this.camera);if(this.showSensorCamera&&this.updateSensorCamera())this.renderSensorCamera(); this.fpsFrames++; let fps=0;if(now-this.lastFpsAt>=500){fps=this.fpsFrames*1000/(now-this.lastFpsAt);this.fpsFrames=0;this.lastFpsAt=now;} const snapshot=this.session&&now-this.lastSnapshotAt>150?(this.lastSnapshotAt=now,this.session.snapshot()):undefined; /* 避免每个 RAF 都触发 Zustand/React 全树重渲染。 */ if(snapshot||fps>0)this.callbacks.onFrame(result,fps,snapshot);}catch(error){this.callbacks.onError(error instanceof Error?error:new Error(String(error)));} this.frame=requestAnimationFrame(this.animate);}; - private updateSensorCamera():boolean {if(!this.session||this.sensorCameraId<0||this.sensorCameraId>=this.session.model.ncam)return false;const id=this.sensorCameraId,p=id*3,m=id*9,data=this.session.data,model=this.session.model;this.sensorCamera.position.set(Number(data.cam_xpos[p]),Number(data.cam_xpos[p+1]),Number(data.cam_xpos[p+2]));const rotation=new THREE.Matrix4().set(Number(data.cam_xmat[m]),Number(data.cam_xmat[m+1]),Number(data.cam_xmat[m+2]),0,Number(data.cam_xmat[m+3]),Number(data.cam_xmat[m+4]),Number(data.cam_xmat[m+5]),0,Number(data.cam_xmat[m+6]),Number(data.cam_xmat[m+7]),Number(data.cam_xmat[m+8]),0,0,0,0,1);this.sensorCamera.quaternion.setFromRotationMatrix(rotation);this.sensorCamera.fov=Number(model.cam_fovy[id])||45;const extent=this.session.geometryBounds().extent;this.sensorCamera.near=Math.max(.001,extent/1000);this.sensorCamera.far=Math.max(100,extent*100);this.sensorCamera.updateProjectionMatrix();return true;} - private renderSensorCamera():void {const size=this.renderer.getSize(new THREE.Vector2()),width=Math.max(120,Math.min(320,size.x*.32)),height=width*9/16,margin=16;this.sensorCamera.aspect=width/height;this.sensorCamera.updateProjectionMatrix();this.renderer.setViewport(margin,margin,width,height);this.renderer.setScissor(margin,margin,width,height);this.renderer.setScissorTest(true);this.renderer.render(this.scene,this.sensorCamera);this.renderer.setScissorTest(false);this.renderer.setViewport(0,0,size.x,size.y);} + private updateSensorCamera():boolean {if(!this.session||this.sensorCameraId<0||this.sensorCameraId>=this.session.model.ncam)return false;const id=this.sensorCameraId,p=id*3,m=id*9,data=this.session.data,model=this.session.model;this.sensorCamera.position.set(Number(data.cam_xpos[p]),Number(data.cam_xpos[p+1]),Number(data.cam_xpos[p+2]));this.sensorRotation.set(Number(data.cam_xmat[m]),Number(data.cam_xmat[m+1]),Number(data.cam_xmat[m+2]),0,Number(data.cam_xmat[m+3]),Number(data.cam_xmat[m+4]),Number(data.cam_xmat[m+5]),0,Number(data.cam_xmat[m+6]),Number(data.cam_xmat[m+7]),Number(data.cam_xmat[m+8]),0,0,0,0,1);this.sensorCamera.quaternion.setFromRotationMatrix(this.sensorRotation);this.sensorCamera.fov=Number(model.cam_fovy[id])||45;const extent=this.session.geometryBounds().extent;this.sensorCamera.near=Math.max(.001,extent/1000);this.sensorCamera.far=Math.max(100,extent*100);this.sensorCamera.updateProjectionMatrix();return true;} + private renderSensorCamera():void {const size=this.renderer.getSize(this.renderSize),width=Math.max(120,Math.min(320,size.x*.32)),height=width*9/16,margin=16;this.sensorCamera.aspect=width/height;this.sensorCamera.updateProjectionMatrix();this.renderer.setViewport(margin,margin,width,height);this.renderer.setScissor(margin,margin,width,height);this.renderer.setScissorTest(true);this.renderer.render(this.scene,this.sensorCamera);this.renderer.setScissorTest(false);this.renderer.setViewport(0,0,size.x,size.y);} - private updateMuJoCoScene():void {const s=this.session!; s.module.mjv_updateScene(s.model,s.data,this.option!,s.perturb,this.mjCamera!,s.module.mjtCatBit.mjCAT_ALL.value,this.mjScene!); const geoms=this.mjScene!.geoms; try {for(let i=0;i=0)return this.meshGeometry(meshIdFromSceneDataId(g.dataid));return new THREE.BufferGeometry();} private meshGeometry(id:number):THREE.BufferGeometry {const m=this.session!.model;const va=Number(m.mesh_vertadr[id]),vn=Number(m.mesh_vertnum[id]),fa=Number(m.mesh_faceadr[id]),fn=Number(m.mesh_facenum[id]);const positions=new Float32Array(vn*3);for(let i=0;i=0){const uv=new Float32Array(tn*2);for(let i=0;i=0,sharedKey=isMesh?`mesh:${meshIdFromSceneDataId(g.dataid)}`:undefined;let geometry=sharedKey?this.geometries.get(sharedKey):undefined;if(!geometry){geometry=this.primitive(g);if(sharedKey)this.geometries.set(sharedKey,geometry);}const map=this.texture(g.texid);const material=new THREE.MeshStandardMaterial({color:new THREE.Color(g.rgba[0],g.rgba[1],g.rgba[2]),opacity:g.rgba[3],transparent:g.rgba[3]<1,...(map?{map}:{}),roughness:Math.max(.05,1-g.shininess),metalness:g.reflectance});const mesh=new THREE.Mesh(geometry,material);mesh.matrixAutoUpdate=false;mesh.castShadow=true;mesh.receiveShadow=true;mesh.userData.geometryKey=key;mesh.userData.ownsGeometry=!sharedKey;return mesh;} private updateMesh(mesh:THREE.Mesh,g:MjvGeom):void {const mat=mesh.material as THREE.MeshStandardMaterial;mat.color.setRGB(g.rgba[0],g.rgba[1],g.rgba[2]);mat.opacity=g.rgba[3];mat.transparent=g.rgba[3]<1;mesh.matrix.set(g.mat[0],g.mat[1],g.mat[2],g.pos[0],g.mat[3],g.mat[4],g.mat[5],g.pos[1],g.mat[6],g.mat[7],g.mat[8],g.pos[2],0,0,0,1);mesh.matrixWorldNeedsUpdate=true;const geomId=g.objtype===this.session!.module.mjtObj.mjOBJ_GEOM.value?g.objid:-1;const bodyId=geomId>=0?Number(this.session!.model.geom_bodyid[geomId]):-1;mesh.userData.geomId=geomId;mesh.userData.bodyId=bodyId;mesh.userData.geomType=g.type;this.applyMeshHighlight(mesh);} private eventPointer(event:PointerEvent):void {const r=this.renderer.domElement.getBoundingClientRect();this.pointer.set((event.clientX-r.left)/r.width*2-1,-((event.clientY-r.top)/r.height)*2+1);} - private onPointerDown=(event:PointerEvent):void=>{this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const hit=this.raycaster.intersectObjects(this.meshes.filter(m=>m.visible),false)[0];if(!hit)return;const mesh=hit.object as THREE.Mesh;this.select(mesh);const bodyId=Number(mesh.userData.bodyId);if(this.mode==='joint'){const joint=this.session?.snapshot().joints.find(j=>j.bodyId===bodyId&&j.editable);if(joint){this.dragStart=this.pointer.clone();this.dragJointId=joint.id;this.dragJointValue=joint.value;const p=joint.bodyId*3,xm=joint.bodyId*9;const origin=new THREE.Vector3(this.session!.data.xpos[p],this.session!.data.xpos[p+1],this.session!.data.xpos[p+2]);const local=new THREE.Vector3(...joint.axis);const axis=new THREE.Vector3(this.session!.data.xmat[xm]*local.x+this.session!.data.xmat[xm+1]*local.y+this.session!.data.xmat[xm+2]*local.z,this.session!.data.xmat[xm+3]*local.x+this.session!.data.xmat[xm+4]*local.y+this.session!.data.xmat[xm+5]*local.z,this.session!.data.xmat[xm+6]*local.x+this.session!.data.xmat[xm+7]*local.y+this.session!.data.xmat[xm+8]*local.z);const a=origin.clone().project(this.camera),b=origin.clone().add(axis).project(this.camera);this.dragAxis.set(b.x-a.x,b.y-a.y);if(this.dragAxis.lengthSq()<1e-6)this.dragAxis.set(1,0);else this.dragAxis.normalize();}}else if(this.mode==='force'&&bodyId>0){this.dragStart=this.pointer.clone();this.session?.initializePerturb(this.mjScene!,bodyId);this.showArrow(hit.point);this.renderer.domElement.setPointerCapture(event.pointerId);}}; - private onPointerMove=(event:PointerEvent):void=>{if(!this.dragStart||!this.session)return;this.eventPointer(event);const dx=this.pointer.x-this.dragStart.x,dy=this.pointer.y-this.dragStart.y;if(this.mode==='joint'&&this.dragJointId>=0)this.session.setJointPosition(this.dragJointId,this.dragJointValue+(dx*this.dragAxis.x+dy*this.dragAxis.y)*Math.PI);else if(this.mode==='force'&&this.selected){const bodyId=Number(this.selected.userData.bodyId),scale=this.forceScale;const force:[number,number,number]=[dx*scale,0,-dy*scale];this.session.setExternalForce(bodyId,force);this.updateArrow(force);}}; + private onPointerDown=(event:PointerEvent):void=>{this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const hit=this.raycaster.intersectObjects(this.meshes.filter(m=>m.visible),false)[0];if(!hit)return;const mesh=hit.object as THREE.Mesh;this.select(mesh);const bodyId=Number(mesh.userData.bodyId);if(this.mode==='joint'){const joint=this.session?.snapshot().joints.find(j=>j.bodyId===bodyId&&j.editable);if(joint){this.dragStart=this.pointer.clone();this.dragJointId=joint.id;this.dragJointType=joint.type;this.dragJointValue=joint.value;this.dragHitDistance=hit.distance;const offset=joint.id*3;this.dragJointPivot.set(Number(this.session!.data.xanchor[offset]),Number(this.session!.data.xanchor[offset+1]),Number(this.session!.data.xanchor[offset+2]));this.dragJointAxisWorld.set(Number(this.session!.data.xaxis[offset]),Number(this.session!.data.xaxis[offset+1]),Number(this.session!.data.xaxis[offset+2])).normalize();this.raycaster.ray.at(hit.distance,this.dragJointStartWorld);const plane=new THREE.Plane().setFromNormalAndCoplanarPoint(this.dragJointAxisWorld,this.dragJointPivot),projected=plane.projectPoint(this.dragJointStartWorld,new THREE.Vector3());this.dragJointStartPlaneVector.copy(projected).sub(this.dragJointPivot);this.dragSlideParameter=closestRayAxisParameter(this.raycaster.ray,this.dragJointPivot,this.dragJointAxisWorld);this.renderer.domElement.setPointerCapture(event.pointerId);}}else if(this.mode==='force'&&bodyId>0){this.dragStart=this.pointer.clone();this.session?.initializePerturb(this.mjScene!,bodyId);this.showArrow(hit.point);this.renderer.domElement.setPointerCapture(event.pointerId);}}; + private onPointerMove=(event:PointerEvent):void=>{if(!this.dragStart||!this.session)return;this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const delta=this.pointer.clone().sub(this.dragStart),rect=this.renderer.domElement.getBoundingClientRect(),aspect=rect.width/Math.max(1,rect.height),screenDelta=new THREE.Vector2(delta.x*aspect,delta.y);if(this.mode==='joint'&&this.dragJointId>=0){let value=this.dragJointValue,currentPlaneVector:THREE.Vector3|undefined,currentSlideParameter=Number.NaN;if(this.dragJointType===2){currentSlideParameter=closestRayAxisParameter(this.raycaster.ray,this.dragJointPivot,this.dragJointAxisWorld);if(Number.isFinite(currentSlideParameter)&&Number.isFinite(this.dragSlideParameter))value+=currentSlideParameter-this.dragSlideParameter;}else if(this.dragJointType===3){const currentWorld=this.raycaster.ray.at(this.dragHitDistance,new THREE.Vector3()),plane=new THREE.Plane().setFromNormalAndCoplanarPoint(this.dragJointAxisWorld,this.dragJointPivot);currentPlaneVector=plane.projectPoint(currentWorld,new THREE.Vector3()).sub(this.dragJointPivot);const worldDelta=signedAngleAroundAxis(this.dragJointStartPlaneVector,currentPlaneVector,this.dragJointAxisWorld),cameraForward=this.camera.getWorldDirection(new THREE.Vector3()),planeFacing=Math.abs(this.raycaster.ray.direction.dot(this.dragJointAxisWorld));const tangentWorld=cameraForward.clone().cross(this.dragJointAxisWorld).normalize(),a=this.dragJointPivot.clone().project(this.camera),b=this.dragJointPivot.clone().add(tangentWorld).project(this.camera),tangentScreen=new THREE.Vector2((b.x-a.x)*aspect,b.y-a.y);const tangentDelta=tangentScreen.lengthSq()>1e-10?screenDelta.dot(tangentScreen.normalize())*Math.PI:0;value+=resolveHingeDragDelta(worldDelta,tangentDelta,planeFacing);}if(this.session.setJointPosition(this.dragJointId,value)){this.dragJointValue=value;this.dragStart.copy(this.pointer);if(currentPlaneVector&¤tPlaneVector.lengthSq()>1e-12)this.dragJointStartPlaneVector.copy(currentPlaneVector);if(Number.isFinite(currentSlideParameter))this.dragSlideParameter=currentSlideParameter;}}else if(this.mode==='force'&&this.selected){const bodyId=Number(this.selected.userData.bodyId),right=new THREE.Vector3(1,0,0).applyQuaternion(this.camera.quaternion),vector=cameraAlignedForce(screenDelta.x,screenDelta.y,right,this.forceScale),force:[number,number,number]=[vector.x,vector.y,vector.z];this.session.setExternalForce(bodyId,force);this.updateArrow(force);}}; private onPointerUp=():void=>{this.stopDrag();}; - private stopDrag():void {this.dragStart=null;this.dragJointId=-1;this.session?.clearExternalForce();if(this.arrow){this.scene.remove(this.arrow);this.arrow.dispose();this.arrow=null;}} + private stopDrag():void {this.dragStart=null;this.dragJointId=-1;this.dragJointType=-1;this.dragHitDistance=0;this.dragSlideParameter=0;this.session?.clearExternalForce();if(this.arrow){this.scene.remove(this.arrow);this.arrow.dispose();this.arrow=null;}} private select(mesh:THREE.Mesh):void {const previous=this.selected;this.selected=mesh;if(previous)this.applyMeshHighlight(previous);this.applyMeshHighlight(mesh);const bodyId=Number(mesh.userData.bodyId),geomId=Number(mesh.userData.geomId);const name=this.session?.snapshot().bodies.find(b=>b.id===bodyId)?.name??`body_${bodyId}`;const e=mesh.matrix.elements;this.callbacks.onSelection({bodyId,geomId,bodyName:name,geomType:Number(mesh.userData.geomType),position:[e[12],e[13],e[14]]});} private showArrow(origin:THREE.Vector3):void {this.arrow=new THREE.ArrowHelper(new THREE.Vector3(1,0,0),origin,0.01,0xf97316);this.scene.add(this.arrow);} private updateArrow(force:[number,number,number]):void {if(!this.arrow)return;const v=new THREE.Vector3(...force),length=v.length()/25;if(length>1e-6){this.arrow.setDirection(v.normalize());this.arrow.setLength(length,Math.min(.2,length*.2),Math.min(.1,length*.1));}} - private disposeMesh(mesh:THREE.Mesh):void {(mesh.material as THREE.Material).dispose();} + private disposeMesh(mesh:THREE.Mesh):void {(mesh.material as THREE.Material).dispose();if(mesh.userData.ownsGeometry)mesh.geometry.dispose();} private releaseModel():void {this.stopDrag();for(const mesh of this.meshes){this.scene.remove(mesh);this.disposeMesh(mesh);}this.meshes=[];for(const g of this.geometries.values())g.dispose();this.geometries.clear();for(const t of this.textures.values())t.dispose();this.textures.clear();this.mjScene?.delete();this.mjCamera?.delete();this.option?.delete();this.mjScene=null;this.mjCamera=null;this.option=null;this.session=null;this.sensorCameraId=-1;this.selected=null;this.highlightedJointId=-1;this.highlightedBodyId=-1;if(this.jointMarker){this.scene.remove(this.jointMarker);this.jointMarker.geometry.dispose();(this.jointMarker.material as THREE.Material).dispose();this.jointMarker=null;}} - dispose():void {cancelAnimationFrame(this.frame);this.releaseModel();this.resizeObserver.disconnect();this.renderer.domElement.removeEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.removeEventListener('pointermove',this.onPointerMove);window.removeEventListener('pointerup',this.onPointerUp);this.controls.dispose();this.orientationGizmo.dispose();this.grid.geometry.dispose();const gridMaterials=Array.isArray(this.grid.material)?this.grid.material:[this.grid.material];for(const material of gridMaterials)material.dispose();this.renderer.dispose();this.renderer.domElement.remove();} + dispose():void {cancelAnimationFrame(this.frame);this.releaseModel();this.resizeObserver.disconnect();this.renderer.domElement.removeEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.removeEventListener('pointermove',this.onPointerMove);this.renderer.domElement.removeEventListener('lostpointercapture',this.onPointerUp);window.removeEventListener('pointerup',this.onPointerUp);window.removeEventListener('pointercancel',this.onPointerUp);window.removeEventListener('blur',this.onPointerUp);this.controls.dispose();this.orientationGizmo.dispose();this.grid.geometry.dispose();const gridMaterials=Array.isArray(this.grid.material)?this.grid.material:[this.grid.material];for(const material of gridMaterials)material.dispose();this.renderer.dispose();this.renderer.domElement.remove();} } diff --git a/wasm/web_platform/src/viewer/interactionMath.test.ts b/wasm/web_platform/src/viewer/interactionMath.test.ts new file mode 100644 index 00000000..eaf7255b --- /dev/null +++ b/wasm/web_platform/src/viewer/interactionMath.test.ts @@ -0,0 +1,27 @@ +import * as THREE from 'three'; +import {cameraAlignedForce,closestRayAxisParameter,resolveHingeDragDelta,signedAngleAroundAxis} from './interactionMath'; + +describe('interactionMath',()=>{ + it('按右手定则计算绕关节轴的旋转方向',()=>{ + expect(signedAngleAroundAxis(new THREE.Vector3(1,0,0),new THREE.Vector3(0,1,0),new THREE.Vector3(0,0,1))).toBeCloseTo(Math.PI/2); + expect(signedAngleAroundAxis(new THREE.Vector3(1,0,0),new THREE.Vector3(0,-1,0),new THREE.Vector3(0,0,1))).toBeCloseTo(-Math.PI/2); + }); + + it('侧视关节平面时采用切线方向',()=>{ + expect(resolveHingeDragDelta(-.4,.25,.05)).toBe(.25); + expect(resolveHingeDragDelta(-.4,.25,.8)).toBe(-.4); + expect(resolveHingeDragDelta(0,.25,.8)).toBe(.25); + }); + + it('通过指针射线和关节轴最近点稳定求解 slide 位移',()=>{ + const axisOrigin=new THREE.Vector3(0,0,0),axis=new THREE.Vector3(1,0,0); + expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(2,0,3),new THREE.Vector3(0,0,-1)),axisOrigin,axis)).toBeCloseTo(2); + expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(-1,0,3),new THREE.Vector3(0,0,-1)),axisOrigin,axis)).toBeCloseTo(-1); + expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(),new THREE.Vector3(1,0,0)),axisOrigin,axis)).toBeNaN(); + }); + + it('外力的右拖和上拖分别映射到相机右方与世界上方',()=>{ + expect(cameraAlignedForce(.5,.25,new THREE.Vector3(0,-1,.3),100).toArray()).toEqual([0,-50,25]); + expect(cameraAlignedForce(.5,0,new THREE.Vector3(-1,0,0),100).x).toBe(-50); + }); +}); diff --git a/wasm/web_platform/src/viewer/interactionMath.ts b/wasm/web_platform/src/viewer/interactionMath.ts new file mode 100644 index 00000000..c13ad5ae --- /dev/null +++ b/wasm/web_platform/src/viewer/interactionMath.ts @@ -0,0 +1,29 @@ +import * as THREE from 'three'; + +/** 计算绕世界轴从 start 到 end 的有符号角度,遵循右手定则。 */ +export function signedAngleAroundAxis(start:THREE.Vector3,end:THREE.Vector3,axis:THREE.Vector3):number { + if(start.lengthSq()<=1e-12||end.lengthSq()<=1e-12||axis.lengthSq()<=1e-12)return Number.NaN; + const a=start.clone().normalize(),b=end.clone().normalize(),normal=axis.clone().normalize(); + return Math.atan2(normal.dot(a.clone().cross(b)),THREE.MathUtils.clamp(a.dot(b),-1,1)); +} + +/** 关节旋转平面接近侧视时,使用相机切线拖动,避免投影退化和方向跳变。 */ +export function resolveHingeDragDelta(worldDelta:number,tangentDelta:number,planeFacingRatio:number,threshold=.2):number { + const tangentValid=Number.isFinite(tangentDelta),worldValid=Number.isFinite(worldDelta)&&(Math.abs(worldDelta)>1e-8||!tangentValid||Math.abs(tangentDelta)<=1e-8); + if(planeFacingRatio