Cen #4
+3
-1
@@ -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",
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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(<LocalTrainingPanel onPolicyReady={vi.fn()}/>);
|
||||
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();
|
||||
});
|
||||
});
|
||||
@@ -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<TrainingServerInfo>();
|
||||
const [job,setJob]=useState<TrainingJob>();
|
||||
const [busy,setBusy]=useState(false),[error,setError]=useState<string>();
|
||||
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<TrainingDevice>('gpu'),[gpuIds,setGpuIds]=useState('0'),[wandbMode,setWandbMode]=useState<WandbMode>('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 <div>
|
||||
<label className="block text-xs text-text-secondary"><span className="mb-1 block">本地训练服务</span><div className="flex gap-2"><input aria-label="本地训练服务地址" className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary" value={endpoint} disabled={active} onChange={event=>setEndpoint(event.target.value)}/><Button icon={<Link className="h-3.5 w-3.5"/>} disabled={busy||active} onClick={()=>void connect()}>连接</Button></div></label>
|
||||
<div className="mt-2 flex items-center justify-between rounded-md border border-border bg-surface px-2 py-1.5 text-[10px] text-text-tertiary"><span className="flex min-w-0 items-center gap-1.5 truncate"><Server className="h-3.5 w-3.5"/>{server?.trainerRoot??'请先启动本地训练服务'}</span><Badge tone={server?.ready?'success':'warning'}>{server?.ready?'可用':'离线'}</Badge></div>
|
||||
{server?.ready&&!job&&<div className="mt-3 space-y-2">
|
||||
<Field label="训练任务"><Select aria-label="训练任务" className="w-full" value={taskId} onChange={event=>setTaskId(event.target.value)}>{server.tasks.map(task=><option key={task} value={task}>{task}</option>)}</Select></Field>
|
||||
<div className="grid grid-cols-2 gap-2"><NumberField label="并行环境" value={numEnvs} min={1} max={16384} onChange={setNumEnvs}/><NumberField label="训练迭代" value={maxIterations} min={1} max={1000000} onChange={setMaxIterations}/><NumberField label="随机种子" value={seed} min={0} max={2147483647} onChange={setSeed}/><Field label="运行名称"><input aria-label="运行名称" className="field h-7 w-full px-2 text-xs text-text-primary" value={runName} onChange={event=>setRunName(event.target.value)}/></Field></div>
|
||||
<div className="grid grid-cols-2 gap-2"><Field label="计算设备"><Select aria-label="计算设备" className="w-full" value={device} onChange={event=>setDevice(event.target.value as TrainingDevice)}><option value="gpu">GPU</option><option value="cpu">CPU</option></Select></Field><Field label="GPU 编号"><input aria-label="GPU 编号" className="field h-7 w-full px-2 text-xs text-text-primary disabled:opacity-40" value={gpuIds} disabled={device==='cpu'} onChange={event=>setGpuIds(event.target.value)}/></Field></div>
|
||||
<Field label="实验记录"><Select aria-label="W&B 模式" className="w-full" value={wandbMode} onChange={event=>setWandbMode(event.target.value as WandbMode)}><option value="offline">本地离线(默认,无需登录)</option><option value="disabled">完全禁用 W&B</option><option value="online">在线 W&B(需要 API Key)</option></Select></Field>
|
||||
<Button variant="primary" className="w-full" icon={<Play className="h-3.5 w-3.5"/>} disabled={busy} onClick={()=>void start()}>发起本地训练</Button>
|
||||
<p className="text-[10px] leading-4 text-text-tertiary">训练使用本地 mjlab 任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。</p>
|
||||
</div>}
|
||||
{job&&<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
|
||||
<div className="mb-2 flex items-center justify-between gap-2"><span className="truncate text-xs font-medium text-text-primary" title={job.id}>{job.taskId}</span><Badge tone={job.state==='succeeded'?'success':job.state==='failed'||job.state==='cancelled'?'warning':'accent'}>{stateLabel(job.state)}</Badge></div>
|
||||
<ProgressBar value={job.progress} label="训练进度"/><div className="mt-2"><PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`}/><PropertyRow label="状态" value={job.message}/></div>
|
||||
{job.logs.length>0&&<details className="mt-2"><summary className="cursor-pointer text-[10px] text-text-secondary">最近日志</summary><pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-[9px] leading-4 text-text-tertiary">{job.logs.slice(-40).join('\n')}</pre></details>}
|
||||
<div className="mt-3 grid grid-cols-2 gap-2">{active?<Button variant="danger" className="col-span-2" icon={<Square className="h-3.5 w-3.5"/>} disabled={busy} onClick={()=>void cancel()}>停止训练</Button>:<><Button disabled={busy||!job.artifactReady} icon={<Download className="h-3.5 w-3.5"/>} onClick={()=>void importResult()}>导入策略</Button><Button onClick={()=>{setJob(undefined);try{localStorage.removeItem(JOB_KEY);}catch{/* ignore */}}}>新建任务</Button></>}</div>
|
||||
</div>}
|
||||
{error&&<p role="alert" className="mt-2 break-words rounded bg-danger/10 p-2 text-[10px] leading-4 text-danger">{error}</p>}
|
||||
</div>;
|
||||
}
|
||||
|
||||
function Field({label,children}:{label:string;children:ReactNode}){return <label className="block text-[10px] text-text-tertiary"><span className="mb-1 block">{label}</span>{children}</label>;}
|
||||
function NumberField({label,value,min,max,onChange}:{label:string;value:number;min:number;max:number;onChange(value:number):void}){return <Field label={label}><input aria-label={label} type="number" className="field h-7 w-full px-2 text-xs text-text-primary" value={value} min={min} max={max} onChange={event=>onChange(Number(event.target.value))}/></Field>;}
|
||||
@@ -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 <ResizablePanel side={side} storageKey={`mujoco-${side}-sidebar-width`} visible={visible}><aside className={`flex h-full w-full min-w-0 flex-col overflow-hidden bg-panel ${side==='left'?'border-r':'border-l'} border-border`}><h2 className="flex h-10 shrink-0 items-center gap-2 border-b border-border bg-panel px-3 text-sm font-semibold text-text-primary"><Settings2 aria-hidden="true" className="h-4 w-4 text-accent"/>{title}</h2>{children}</aside></ResizablePanel>;}
|
||||
|
||||
@@ -33,7 +34,7 @@ export function ModelControlsSidebar(props:ModelControlsProps){const [tab,setTab
|
||||
const properties=<><CollapsibleSection title="模型信息" defaultOpen badge={<Badge>{s.model.nbody} Body</Badge>}><div><PropertyRow label="Body" value={s.model.nbody}/><PropertyRow label="Joint" value={s.model.njnt}/><PropertyRow label="Geom" value={s.model.ngeom}/><PropertyRow label="Actuator" value={s.model.nactuator}/><PropertyRow label="qpos / qvel" value={`${s.model.nq} / ${s.model.nv}`}/></div></CollapsibleSection>
|
||||
{props.selectedFormat==='urdf'&&<CollapsibleSection title="URDF 处理方式" defaultOpen={false}><Select aria-label="URDF 处理方式" className="w-full" value={props.urdfMode} disabled={props.loading} onChange={event=>props.onUrdfMode(event.target.value as UrdfLoadMode)}><option value="mjcf">转换为 MJCF(推荐)</option><option value="native">MuJoCo 原生 URDF</option></Select><label className="mt-3 block text-xs text-text-secondary"><span className="mb-1 block">基座类型</span><Select aria-label="URDF 基座类型" className="w-full" value={props.baseMode} disabled={props.loading||props.urdfMode==='native'} onChange={event=>props.onBaseMode(event.target.value as UrdfBaseMode)}><option value="floating">浮动基座(Free Joint)</option><option value="fixed">固定基座(连接世界)</option></Select></label><p className="mt-2 text-xs text-text-tertiary">MJCF 模式保留 visual mesh、添加物理地面,并将模型最低点对齐到 z=0。</p><Check label="显示碰撞几何" checked={props.showCollision} onChange={props.onShowCollision}/></CollapsibleSection>}
|
||||
<CollapsibleSection title="当前选择" defaultOpen>{props.selection?<div className="text-xs"><PropertyRow label="Body" value={props.selection.bodyName} action={<CopyButton value={props.selection.bodyName} label="复制 Body 名称"/>}/><PropertyRow label="标识" value={`${props.selection.bodyId} / ${props.selection.geomId} / ${props.selection.geomType}`} action={<CopyButton value={`body ${props.selection.bodyId}, geom ${props.selection.geomId}, type ${props.selection.geomType}`} label="复制标识"/>}/><PropertyRow label="位置" value={props.selection.position.map(value=>value.toFixed(3)).join(', ')} action={<CopyButton value={props.selection.position.join(', ')} label="复制位置"/>}/></div>:<p className="flex items-center gap-2 text-xs text-text-tertiary"><Info className="h-3.5 w-3.5"/>在视口中单击物体</p>}</CollapsibleSection></>;
|
||||
const controls=<><CollapsibleSection title="ONNX 强化学习策略" defaultOpen badge={s.rlPolicy?<Badge>{s.rlPolicy.enabled?'推理':'停止'}</Badge>:undefined}><RLPolicyPanel paths={props.policyPaths} selectedPath={props.selectedPolicyPath} status={props.policyStatus??s.rlPolicy} loading={props.loading} onSelectPath={props.onSelectPolicyPath} onLoadPath={props.onLoadPolicyPath} onImport={props.onImportPolicy} onToggle={props.onTogglePolicy} onCommand={props.onPolicyCommand} onRemove={props.onRemovePolicy}/></CollapsibleSection><CollapsibleSection title="Python 控制器" defaultOpen badge={s.controller?<Badge>{s.controller.enabled?'运行':'停止'}</Badge>:undefined}><PythonControllerPanel paths={props.controllerPaths} selectedPath={props.selectedControllerPath} status={props.controllerStatus??s.controller} loading={props.loading} onSelectPath={props.onSelectControllerPath} onLoadPath={props.onLoadControllerPath} onImport={props.onImportController} onToggle={props.onToggleController} onCommand={props.onControllerCommand} onRemove={props.onRemoveController}/></CollapsibleSection><CollapsibleSection title="Actuator" defaultOpen={false} badge={<Badge>{s.actuators.length}</Badge>}>{s.actuators.length?s.actuators.map(actuator=><ActuatorControl key={actuator.id} actuator={actuator} onControl={value=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):<p className="text-xs text-text-tertiary">模型没有驱动器</p>}</CollapsibleSection>
|
||||
const controls=<><CollapsibleSection title="ONNX 强化学习策略" defaultOpen badge={s.rlPolicy?<Badge>{s.rlPolicy.enabled?'推理':'停止'}</Badge>:undefined}><RLPolicyPanel paths={props.policyPaths} selectedPath={props.selectedPolicyPath} status={props.policyStatus??s.rlPolicy} loading={props.loading} onSelectPath={props.onSelectPolicyPath} onLoadPath={props.onLoadPolicyPath} onImport={props.onImportPolicy} onToggle={props.onTogglePolicy} onCommand={props.onPolicyCommand} onRemove={props.onRemovePolicy}/></CollapsibleSection><CollapsibleSection title="本地强化学习训练" defaultOpen={false}><LocalTrainingPanel onPolicyReady={props.onImportPolicy}/></CollapsibleSection><CollapsibleSection title="Python 控制器" defaultOpen badge={s.controller?<Badge>{s.controller.enabled?'运行':'停止'}</Badge>:undefined}><PythonControllerPanel paths={props.controllerPaths} selectedPath={props.selectedControllerPath} status={props.controllerStatus??s.controller} loading={props.loading} onSelectPath={props.onSelectControllerPath} onLoadPath={props.onLoadControllerPath} onImport={props.onImportController} onToggle={props.onToggleController} onCommand={props.onControllerCommand} onRemove={props.onRemoveController}/></CollapsibleSection><CollapsibleSection title="Actuator" defaultOpen={false} badge={<Badge>{s.actuators.length}</Badge>}>{s.actuators.length?s.actuators.map(actuator=><ActuatorControl key={actuator.id} actuator={actuator} onControl={value=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):<p className="text-xs text-text-tertiary">模型没有驱动器</p>}</CollapsibleSection>
|
||||
<CollapsibleSection title="关节" defaultOpen badge={<Badge>{s.joints.length}</Badge>}><div className="mb-4 grid grid-cols-2 gap-2"><Button onClick={props.onResetJoints}>重置关节</Button><Button variant={props.ignoreJointLimits?'primary':'secondary'} aria-pressed={props.ignoreJointLimits} onClick={props.onToggleJointLimits}>忽略关节限位</Button><Button variant={props.jointAdvanced?'primary':'secondary'} aria-pressed={props.jointAdvanced} onClick={props.onToggleAdvanced}>高级</Button><Button variant={props.angleUnit==='deg'?'primary':'secondary'} aria-pressed={props.angleUnit==='deg'} onClick={props.onToggleAngleUnit}>{props.angleUnit==='rad'?'rad 弧度制':'° 角度制'}</Button></div>{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 <ControlSlider key={joint.id} label={`${joint.name}${joint.editable?'':'(只读)'}`} value={joint.value*scale} min={joint.min*scale} max={joint.max*scale} unit={unit} advanced={props.jointAdvanced} limited={joint.limited} limitsIgnored={joint.limitsIgnored} limitMin={joint.limitMin*scale} limitMax={joint.limitMax*scale} disabled={!joint.editable} onChange={value=>props.onJoint(joint.id,value/scale)}/>;})}</CollapsibleSection>
|
||||
<CollapsibleSection title="外力强度" defaultOpen={false}><ControlSlider label={`${props.forceScale.toFixed(0)} N/屏幕单位`} value={props.forceScale} min={5} max={200} onChange={props.onForceScale}/><p className="text-xs text-text-tertiary">选择“外力施加”,在动态物体上按住拖动,松开即清零。</p></CollapsibleSection></>;
|
||||
return <SidebarPanel title="模型与控制" side="right" visible={props.visible}><Tabs label="模型控制侧栏" value={tab} onValueChange={setTab} items={[{value:'properties',label:'属性',icon:<Info className="h-3.5 w-3.5"/>,content:properties},{value:'controls',label:'控制',icon:<SlidersHorizontal className="h-3.5 w-3.5"/>,content:controls}]}/></SidebarPanel>;
|
||||
|
||||
@@ -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');});
|
||||
});
|
||||
@@ -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<Error>{
|
||||
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<T>(path:string,init?:RequestInit):Promise<T>{
|
||||
const response=await fetch(`${this.endpoint}${path}`,init);
|
||||
if(!response.ok)throw await responseError(response);
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
health():Promise<TrainingServerInfo>{return this.json('/api/training/health');}
|
||||
start(request:TrainingRequest):Promise<TrainingJob>{return this.json('/api/training/jobs',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify(request)});}
|
||||
job(id:string):Promise<TrainingJob>{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`);}
|
||||
cancel(id:string):Promise<TrainingJob>{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`,{method:'DELETE'});}
|
||||
async downloadPolicy(id:string):Promise<File>{
|
||||
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'});
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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<string,THREE.BufferGeometry>(); private textures=new Map<number,THREE.DataTexture>();
|
||||
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<geoms.size();i++){const geom=geoms.get(i);if(!geom)continue;try{let mesh=this.meshes[i];const key=this.geometryKey(geom);if(!mesh||mesh.userData.geometryKey!==key){if(mesh){this.scene.remove(mesh);this.disposeMesh(mesh);} mesh=this.createMesh(geom,key);this.meshes[i]=mesh;this.scene.add(mesh);} mesh.visible=true;this.updateMesh(mesh,geom);}finally{geom.delete();}} for(let i=geoms.size();i<this.meshes.length;i++)this.meshes[i].visible=false;}finally{geoms.delete();}}
|
||||
private geometryKey(g:MjvGeom):string {const dataId=g.type===this.session!.module.mjtGeom.mjGEOM_MESH.value?meshIdFromSceneDataId(g.dataid):g.dataid;return `${g.type}:${dataId}:${Array.from(g.size).join(',')}`;}
|
||||
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 {const count=geoms.size();for(let i=0;i<count;i++){const geom=geoms.get(i);if(!geom)continue;try{let mesh=this.meshes[i];const key=this.geometryKey(geom);if(!mesh||mesh.userData.geometryKey!==key){if(mesh){this.scene.remove(mesh);this.disposeMesh(mesh);} mesh=this.createMesh(geom,key);this.meshes[i]=mesh;this.scene.add(mesh);} mesh.visible=true;this.updateMesh(mesh,geom);}finally{geom.delete();}} for(let i=count;i<this.meshes.length;i++)this.meshes[i].visible=false;}finally{geoms.delete();}}
|
||||
/** 模型 geom 的尺寸固定,以 objid 保持稳定;接触点等动态 geom 才包含尺寸。 */
|
||||
private geometryKey(g:MjvGeom):string {const m=this.session!.module,dataId=g.type===m.mjtGeom.mjGEOM_MESH.value?meshIdFromSceneDataId(g.dataid):g.dataid;if(g.objtype===m.mjtObj.mjOBJ_GEOM.value)return `model:${g.objid}:${g.type}:${dataId}`;return `dynamic:${g.type}:${dataId}:${Array.from(g.size).join(',')}`;}
|
||||
private primitive(g:MjvGeom):THREE.BufferGeometry {const m=this.session!.module,t=g.type,s=g.size;if(t===m.mjtGeom.mjGEOM_PLANE.value)return new THREE.PlaneGeometry(2*(s[0]||1e3),2*(s[1]||1e3));if(t===m.mjtGeom.mjGEOM_SPHERE.value)return new THREE.SphereGeometry(s[0],24,16);if(t===m.mjtGeom.mjGEOM_CAPSULE.value)return new CapsuleGeometry(s[0],2*s[2]);if(t===m.mjtGeom.mjGEOM_BOX.value)return new THREE.BoxGeometry(2*s[0],2*s[1],2*s[2]);if(t===m.mjtGeom.mjGEOM_CYLINDER.value){const x=new THREE.CylinderGeometry(s[0],s[0],2*s[2],24);x.rotateX(Math.PI/2);return x;}if(t===m.mjtGeom.mjGEOM_ELLIPSOID.value){const x=new THREE.SphereGeometry(1,24,16);x.scale(s[0],s[1],s[2]);return x;}if(t===m.mjtGeom.mjGEOM_MESH.value&&g.dataid>=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<positions.length;i++)positions[i]=m.mesh_vert[va*3+i];const indices=new Uint32Array(fn*3);for(let i=0;i<indices.length;i++)indices[i]=m.mesh_face[fa*3+i];const geometry=new THREE.BufferGeometry();geometry.setAttribute('position',new THREE.BufferAttribute(positions,3));geometry.setIndex(new THREE.BufferAttribute(indices,1));const na=Number(m.mesh_normaladr[id]),nn=Number(m.mesh_normalnum[id]);if(nn===vn){const normals=new Float32Array(nn*3);for(let i=0;i<normals.length;i++)normals[i]=m.mesh_normal[na*3+i];geometry.setAttribute('normal',new THREE.BufferAttribute(normals,3));}else geometry.computeVertexNormals();const ta=Number(m.mesh_texcoordadr[id]),tn=Number(m.mesh_texcoordnum[id]);if(tn===vn&&ta>=0){const uv=new Float32Array(tn*2);for(let i=0;i<uv.length;i++)uv[i]=m.mesh_texcoord[ta*2+i];geometry.setAttribute('uv',new THREE.BufferAttribute(uv,2));}geometry.computeBoundingSphere();return geometry;}
|
||||
private texture(id:number):THREE.DataTexture|undefined {if(id<0)return;let found=this.textures.get(id);if(found)return found;const m=this.session!.model,w=Number(m.tex_width[id]),h=Number(m.tex_height[id]),channels=Number(m.tex_nchannel[id]||3),adr=Number(m.tex_adr[id]);if(!w||!h)return;const data=new Uint8Array(w*h*channels);for(let i=0;i<data.length;i++)data[i]=m.tex_data[adr+i];found=new THREE.DataTexture(data,w,h,channels===4?THREE.RGBAFormat:THREE.RGBFormat);found.colorSpace=THREE.SRGBColorSpace;found.flipY=true;found.needsUpdate=true;this.textures.set(id,found);return found;}
|
||||
private createMesh(g:MjvGeom,key:string):THREE.Mesh {let geometry=this.geometries.get(key);if(!geometry){geometry=this.primitive(g);this.geometries.set(key,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;return mesh;}
|
||||
// 只缓存真正可复用的模型 mesh。动态接触几何会逐帧改变尺寸,必须由所属
|
||||
// THREE.Mesh 在替换时释放,否则 geometry key 会无限增长并最终耗尽标签页内存。
|
||||
private createMesh(g:MjvGeom,key:string):THREE.Mesh {const m=this.session!.module,isMesh=g.type===m.mjtGeom.mjGEOM_MESH.value&&g.dataid>=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();}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
@@ -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<threshold&&tangentValid)return tangentDelta;
|
||||
if(worldValid)return worldDelta;
|
||||
return tangentValid?tangentDelta:0;
|
||||
}
|
||||
|
||||
/** 求指针射线和世界关节轴两条直线的最近点参数,参数单位与 slide qpos 一致。 */
|
||||
export function closestRayAxisParameter(ray:THREE.Ray,axisOrigin:THREE.Vector3,axisDirection:THREE.Vector3):number {
|
||||
const direction=ray.direction.clone().normalize(),axis=axisDirection.clone().normalize();if(direction.lengthSq()<=1e-12||axis.lengthSq()<=1e-12)return Number.NaN;
|
||||
const offset=ray.origin.clone().sub(axisOrigin),dot=direction.dot(axis),denominator=1-dot*dot;if(denominator<=1e-8)return Number.NaN;
|
||||
const value=(axis.dot(offset)-dot*direction.dot(offset))/denominator;return Number.isFinite(value)?value:Number.NaN;
|
||||
}
|
||||
|
||||
/** MuJoCo MOVE_V 风格:水平跟随相机右方向,垂直始终对应世界 +Z。 */
|
||||
export function cameraAlignedForce(dx:number,dy:number,cameraRight:THREE.Vector3,scale:number):THREE.Vector3 {
|
||||
const right=cameraRight.clone();right.z=0;if(right.lengthSq()<=1e-12)right.set(1,0,0);else right.normalize();
|
||||
return right.multiplyScalar(dx*scale).addScaledVector(new THREE.Vector3(0,0,1),dy*scale);
|
||||
}
|
||||
Reference in New Issue
Block a user