feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
This commit is contained in:
+273
-13
@@ -25,6 +25,9 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, unquote, urlsplit
|
||||
|
||||
from mobile_manipulator.config import MOBILE_TASKS, validate_mobile_params
|
||||
from mobile_manipulator.config import TASK as MOBILE_CONTRACT
|
||||
from mobile_manipulator.packages import MAX_UPLOAD, MobilePackages
|
||||
from pretrained_sources import PretrainedSources, SourceError
|
||||
from task_config import (
|
||||
OBSTACLE_TASK,
|
||||
@@ -38,10 +41,10 @@ from tuning.process import GpuLease, ResourceBusyError
|
||||
from tuning.schema import RewardConfigError, validate_configuration
|
||||
from tuning.scoring import EvaluationError
|
||||
|
||||
VERSION = "0.4.0"
|
||||
VERSION = "0.6.0"
|
||||
MAX_REQUEST_BYTES = 128 * 1024 # Bounded full boxes-v1 payload (<=257 boxes).
|
||||
# Rough 可以训练,但其高度扫描 actor 不允许冒充浏览器 Flat 部署。
|
||||
DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK)
|
||||
DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK, *MOBILE_TASKS)
|
||||
ACTIVE_STATES = {"queued", "running"}
|
||||
MAX_JOBS = 20
|
||||
ANSI_ESCAPE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]")
|
||||
@@ -86,6 +89,9 @@ class TrainingConfig:
|
||||
task_config: dict[str, Any] | None = None
|
||||
deployment: dict[str, Any] = field(default_factory=dict)
|
||||
pretrained: dict[str, Any] | None = None
|
||||
mobile_package_id: str | None = None
|
||||
mobile_params: dict[str, Any] | None = None
|
||||
mobile_checkpoint: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -135,9 +141,11 @@ class TrainingManager:
|
||||
check_environment: bool = True,
|
||||
lease: GpuLease | None = None,
|
||||
sources: PretrainedSources | None = None,
|
||||
mobile_python: str | None = None,
|
||||
):
|
||||
self.trainer_root = trainer_root.expanduser().resolve()
|
||||
self.python = str(Path(python).expanduser()) if os.sep in python else python
|
||||
self.mobile_python = mobile_python or self.python
|
||||
self.tasks = tasks
|
||||
self.jobs: dict[str, TrainingJob] = {}
|
||||
self.lock = threading.RLock()
|
||||
@@ -146,8 +154,39 @@ class TrainingManager:
|
||||
self.lease = lease or GpuLease()
|
||||
self.preset_resolver: Any = None
|
||||
self.sources = sources
|
||||
self.mobile_packages = MobilePackages(self.trainer_root / "logs" / "mobile_packages")
|
||||
self._mobile_environment_error: str | None | bool = False
|
||||
|
||||
def readiness_error(self) -> str | None:
|
||||
def readiness_error(self, task_id: str | None = None) -> str | None:
|
||||
if task_id in MOBILE_TASKS:
|
||||
if not self.check_environment:
|
||||
return None
|
||||
if self._mobile_environment_error is False:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
self.mobile_python,
|
||||
"-c",
|
||||
(
|
||||
"import mujoco,gymnasium,torch,stable_baselines3,onnx,onnxruntime; "
|
||||
"assert mujoco.__version__ == '3.11.0', "
|
||||
"'移动操作需要 MuJoCo 3.11.0,请配置 --mobile-python;"
|
||||
"不要升级 Go2 环境'"
|
||||
),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
check=False,
|
||||
)
|
||||
self._mobile_environment_error = (
|
||||
"移动操作 Python 依赖不可用:" + result.stderr[-1500:]
|
||||
if result.returncode
|
||||
else None
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired) as error:
|
||||
self._mobile_environment_error = f"无法检查移动操作环境:{error}"
|
||||
return self._mobile_environment_error or None
|
||||
if not self.trainer_root.is_dir():
|
||||
return f"训练工程目录不存在:{self.trainer_root}"
|
||||
if not (self.trainer_root / "scripts" / "train.py").is_file():
|
||||
@@ -184,20 +223,26 @@ class TrainingManager:
|
||||
return next((job.id for job in self.jobs.values() if job.state in ACTIVE_STATES), None)
|
||||
|
||||
def health(self) -> dict[str, Any]:
|
||||
error = self.readiness_error()
|
||||
errors = {task: self.readiness_error(task) for task in self.tasks}
|
||||
ready = any(error is None for error in errors.values())
|
||||
error = None if ready else next(iter(errors.values()), "没有可用任务")
|
||||
metadata = task_metadata(self.tasks)
|
||||
for item in metadata:
|
||||
item.update(ready=errors[item["id"]] is None, error=errors[item["id"]])
|
||||
return {
|
||||
"version": VERSION,
|
||||
"ready": error is None,
|
||||
"ready": ready,
|
||||
"trainerRoot": str(self.trainer_root),
|
||||
"python": self.python,
|
||||
"tasks": list(self.tasks),
|
||||
"pretrainedSources": self.sources.catalog() if self.sources else [],
|
||||
"pretrainedUpload": {
|
||||
"enabled": self.sources is not None, "templateId": "go2-legacy47-v1",
|
||||
"enabled": self.sources is not None,
|
||||
"templateId": "go2-legacy47-v1",
|
||||
"formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2},
|
||||
"endpoint": "/api/training/pretrained-sources/upload",
|
||||
},
|
||||
"taskMetadata": task_metadata(self.tasks),
|
||||
"taskMetadata": metadata,
|
||||
"activeJobId": self.active_job_id(),
|
||||
"error": error,
|
||||
}
|
||||
@@ -221,6 +266,8 @@ class TrainingManager:
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
"mobilePackageId",
|
||||
"mobileParams",
|
||||
}
|
||||
if payload.keys() - allowed:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "请求包含未知字段(不接受配置路径/MJCF)")
|
||||
@@ -276,6 +323,56 @@ class TrainingManager:
|
||||
except RewardConfigError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
seed = integer("seed", 0, 2_147_483_647)
|
||||
if task_id in MOBILE_TASKS:
|
||||
if any(
|
||||
k in payload
|
||||
for k in (
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
"pretrainedSourceId",
|
||||
)
|
||||
):
|
||||
raise ApiError(
|
||||
HTTPStatus.BAD_REQUEST, "移动操作任务不接受 Go2 地形/传感器/预训练参数"
|
||||
)
|
||||
try:
|
||||
package_id = payload.get("mobilePackageId")
|
||||
package = self.mobile_packages.describe(package_id)
|
||||
if package["robotId"] != MOBILE_TASKS[task_id]:
|
||||
raise ValueError("训练任务与场景机器人变体不匹配")
|
||||
params = validate_mobile_params(payload.get("mobileParams", {}))
|
||||
checkpoint = self.mobile_resume(params, task_id, package, package_id)
|
||||
except (ValueError, OSError) as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
if device == "gpu" and len(raw_gpu_ids) != 1:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "移动操作 PPO 仅支持单个 GPU")
|
||||
return TrainingConfig(
|
||||
task_id=task_id,
|
||||
num_envs=integer("numEnvs", 1, 64),
|
||||
max_iterations=integer("maxIterations", 1, 1_000_000),
|
||||
seed=seed,
|
||||
run_name=run_name,
|
||||
device=device,
|
||||
gpu_ids=raw_gpu_ids,
|
||||
wandb_mode=wandb_mode,
|
||||
mobile_package_id=package_id,
|
||||
mobile_params=params,
|
||||
mobile_checkpoint=checkpoint,
|
||||
deployment={
|
||||
"trainingStage": params["stage"],
|
||||
"actionSemantics": MOBILE_CONTRACT["actionSemantics"],
|
||||
"version": 1,
|
||||
"browserCompatible": True,
|
||||
"trainingTaskId": task_id,
|
||||
"taskId": MOBILE_CONTRACT["id"],
|
||||
**package,
|
||||
},
|
||||
)
|
||||
if "mobilePackageId" in payload or "mobileParams" in payload:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "Go2 任务不接受移动操作参数")
|
||||
try:
|
||||
custom = validate_task_config(task_id, payload, seed)
|
||||
except TaskConfigError as error:
|
||||
@@ -308,11 +405,70 @@ class TrainingManager:
|
||||
reward_config=reward_config,
|
||||
)
|
||||
|
||||
def mobile_resume(self, params, task_id, package, package_id) -> str | None:
|
||||
source_id = params.get("sourceJobId")
|
||||
if source_id is None:
|
||||
if params["stage"] != "navigate":
|
||||
raise ValueError("请先完成底盘接近训练,再选择通过评估的前一阶段作业")
|
||||
return None
|
||||
with self.lock:
|
||||
source = self.jobs.get(source_id)
|
||||
if (
|
||||
not source
|
||||
or source.state != "succeeded"
|
||||
or source.config.task_id != task_id
|
||||
or not source.artifact
|
||||
):
|
||||
raise ValueError("接续训练需要同一机器人已完成的服务内作业")
|
||||
if source.config.mobile_package_id != package_id:
|
||||
raise ValueError("接续场景/资产快照不匹配;资产变化后请重新训练")
|
||||
deployment = source.config.deployment
|
||||
if (
|
||||
any(
|
||||
deployment.get(key) != package.get(key)
|
||||
for key in ("robotId", "sceneSha256", "robotConfigSha256")
|
||||
)
|
||||
or deployment.get("actionSemantics") != MOBILE_CONTRACT["actionSemantics"]
|
||||
or deployment.get("taskId") != MOBILE_CONTRACT["id"]
|
||||
):
|
||||
raise ValueError("接续作业与当前场景/安全控制契约不匹配")
|
||||
stages = ("navigate", "reach", "pick-place")
|
||||
previous = deployment.get("trainingStage")
|
||||
if (
|
||||
previous not in stages
|
||||
or not stages.index(previous)
|
||||
<= stages.index(params["stage"])
|
||||
<= stages.index(previous) + 1
|
||||
):
|
||||
raise ValueError("只支持同阶段续训或依次推进:底盘接近 → 末端接近 → 抓取放置")
|
||||
if previous != params["stage"]:
|
||||
if any(
|
||||
params[key] != (source.config.mobile_params or {}).get(key)
|
||||
for key in ("objectPosition", "goalPosition", "positionJitter")
|
||||
):
|
||||
raise ValueError(
|
||||
"升级阶段必须保留已评估的初态分布;改变坐标/随机范围请先同阶段续训"
|
||||
)
|
||||
evaluation = deployment.get("evaluation", {})
|
||||
if (
|
||||
evaluation.get("episodes", 0) < 10
|
||||
or evaluation.get("successRate", 0) < MOBILE_CONTRACT["navigationSuccessRate"]
|
||||
or evaluation.get("safetyStops", 1) != 0
|
||||
):
|
||||
raise ValueError(
|
||||
"上一阶段尚未达标:至少 10 回合独立评估、成功率 ≥80%、"
|
||||
"无安全终止;请先同阶段续训"
|
||||
)
|
||||
checkpoint = source.artifact.with_suffix(".ppo.zip")
|
||||
if not checkpoint.is_file():
|
||||
raise ValueError("接续作业缺少 PPO checkpoint")
|
||||
return str(checkpoint)
|
||||
|
||||
def start(self, payload: Any) -> dict[str, Any]:
|
||||
error = self.readiness_error()
|
||||
config = self.parse_config(payload)
|
||||
error = self.readiness_error(config.task_id)
|
||||
if error:
|
||||
raise ApiError(HTTPStatus.SERVICE_UNAVAILABLE, error)
|
||||
config = self.parse_config(payload)
|
||||
with self.lock:
|
||||
if self.active_job_id():
|
||||
raise ApiError(HTTPStatus.CONFLICT, "已有训练任务正在运行,请等待完成或先停止任务")
|
||||
@@ -357,6 +513,12 @@ class TrainingManager:
|
||||
raise ApiError(HTTPStatus.NOT_FOUND, "该训练任务尚未生成 policy.onnx")
|
||||
return job.artifact
|
||||
|
||||
def deployment_artifact(self, job_id: str) -> Path:
|
||||
artifact = self.artifact(job_id).with_name("deployment.json")
|
||||
if not artifact.is_file():
|
||||
raise ApiError(HTTPStatus.NOT_FOUND, "该任务没有独立部署元数据")
|
||||
return artifact
|
||||
|
||||
def cancel(self, job_id: str) -> dict[str, Any]:
|
||||
with self.lock:
|
||||
job = self.jobs.get(job_id)
|
||||
@@ -400,6 +562,27 @@ class TrainingManager:
|
||||
def command_for(
|
||||
self, config: TrainingConfig, task_config_path: Path | None = None
|
||||
) -> list[str]:
|
||||
if config.task_id in MOBILE_TASKS:
|
||||
return [
|
||||
self.mobile_python,
|
||||
"-u",
|
||||
"-m",
|
||||
"training_server.mobile_manipulator.train",
|
||||
"--package",
|
||||
str(self.mobile_packages.path(config.mobile_package_id)),
|
||||
"--iterations",
|
||||
str(config.max_iterations),
|
||||
"--num-envs",
|
||||
str(config.num_envs),
|
||||
"--seed",
|
||||
str(config.seed),
|
||||
"--device",
|
||||
"cpu" if config.device == "cpu" else f"cuda:{config.gpu_ids[0]}",
|
||||
"--params",
|
||||
json.dumps(config.mobile_params),
|
||||
"--task-id",
|
||||
config.task_id,
|
||||
] + (["--resume", config.mobile_checkpoint] if config.mobile_checkpoint else [])
|
||||
command = [
|
||||
self.python,
|
||||
"-u",
|
||||
@@ -480,6 +663,26 @@ class TrainingManager:
|
||||
config_path = job_dir / "training_config.json"
|
||||
config_path.write_text(json.dumps(job.config.task_config), encoding="utf-8")
|
||||
command = self.command_for(job.config, config_path)
|
||||
mobile = job.config.task_id in MOBILE_TASKS
|
||||
if mobile:
|
||||
job_dir = self.trainer_root / "logs" / "rsl_rl" / "web_jobs" / job.id
|
||||
job_dir.mkdir(parents=True, exist_ok=True)
|
||||
command.extend(("--output", str(job_dir / "policy.onnx")))
|
||||
(job_dir / "training_config.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"taskId": job.config.task_id,
|
||||
"mobilePackageId": job.config.mobile_package_id,
|
||||
"mobileParams": job.config.mobile_params,
|
||||
"numEnvs": job.config.num_envs,
|
||||
"maxIterations": job.config.max_iterations,
|
||||
"seed": job.config.seed,
|
||||
"device": job.config.device,
|
||||
"gpuIds": job.config.gpu_ids,
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
if config_path is not None:
|
||||
command.extend(("--output-dir", str(config_path.parent)))
|
||||
# Popen 与 process 登记必须和取消检查处于同一个临界区:cancel() 要么在
|
||||
@@ -490,7 +693,7 @@ class TrainingManager:
|
||||
return
|
||||
process = subprocess.Popen(
|
||||
command,
|
||||
cwd=self.trainer_root,
|
||||
cwd=Path(__file__).resolve().parent.parent if mobile else self.trainer_root,
|
||||
env=environment,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
@@ -513,7 +716,22 @@ class TrainingManager:
|
||||
finally:
|
||||
process.stdout.close()
|
||||
return_code = process.wait()
|
||||
artifact = self._find_artifact(before)
|
||||
artifact = (job_dir / "policy.onnx") if mobile else self._find_artifact(before)
|
||||
if artifact is not None and not artifact.is_file():
|
||||
artifact = None
|
||||
if mobile and return_code == 0 and not job.cancel_requested and artifact:
|
||||
deployment = json.loads((job_dir / "deployment.json").read_text())
|
||||
for key in (
|
||||
"trainingTaskId",
|
||||
"robotId",
|
||||
"sceneSha256",
|
||||
"robotConfigSha256",
|
||||
"trainingStage",
|
||||
"actionSemantics",
|
||||
):
|
||||
if deployment.get(key) != job.config.deployment.get(key):
|
||||
raise ValueError(f"导出部署元数据不匹配:{key}")
|
||||
job.config.deployment = deployment
|
||||
with self.lock:
|
||||
job.process = None
|
||||
job.ended_at = now_iso()
|
||||
@@ -543,7 +761,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
tuning_manager: TuningManager
|
||||
allowed_origins: tuple[str, ...] = ()
|
||||
access_token = ""
|
||||
server_version = "MuJoCoLocalTraining/0.4"
|
||||
server_version = "MuJoCoLocalTraining/0.6"
|
||||
|
||||
def log_message(self, format: str, *args: Any) -> None:
|
||||
sys.stderr.write(f"[{self.log_date_time_string()}] {format % args}\n")
|
||||
@@ -644,7 +862,12 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传显示名称过长")
|
||||
try:
|
||||
result = self.manager.sources.receive_upload(
|
||||
self.rfile, length, fmt, template, name, set_timeout=self.connection.settimeout,
|
||||
self.rfile,
|
||||
length,
|
||||
fmt,
|
||||
template,
|
||||
name,
|
||||
set_timeout=self.connection.settimeout,
|
||||
)
|
||||
except OSError as error:
|
||||
raise ApiError(
|
||||
@@ -734,6 +957,14 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
if match:
|
||||
self._send_file(self.tuning_manager.best_artifact(match.group(1)), "policy.onnx")
|
||||
return
|
||||
deployment_match = re.fullmatch(
|
||||
r"/api/training/jobs/([0-9a-f]{32})/artifacts/deployment\.json", path
|
||||
)
|
||||
if deployment_match:
|
||||
self._send_file(
|
||||
self.manager.deployment_artifact(deployment_match.group(1)), "deployment.json"
|
||||
)
|
||||
return
|
||||
job_id, artifact = self._route(path)
|
||||
if not job_id:
|
||||
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
|
||||
@@ -751,6 +982,29 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
if path == "/api/training/pretrained-sources/upload":
|
||||
self._upload()
|
||||
return
|
||||
if path == "/api/training/mobile-packages":
|
||||
self.close_connection = True
|
||||
lengths = self.headers.get_all("Content-Length", [])
|
||||
if (
|
||||
self.headers.get("Transfer-Encoding")
|
||||
or self.headers.get("Content-Encoding")
|
||||
or self.headers.get("Content-Type") != "application/zip"
|
||||
or len(lengths) != 1
|
||||
or not re.fullmatch(r"[0-9]{1,10}", lengths[0])
|
||||
):
|
||||
raise ApiError(
|
||||
HTTPStatus.BAD_REQUEST, "场景上传需要 application/zip 和唯一 Content-Length"
|
||||
)
|
||||
length = int(lengths[0])
|
||||
if not 0 < length <= MAX_UPLOAD:
|
||||
raise ApiError(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "场景上传上限 128 MiB")
|
||||
self.connection.settimeout(60)
|
||||
try:
|
||||
result = self.manager.mobile_packages.receive(self.rfile, length)
|
||||
except ValueError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
self._json(HTTPStatus.CREATED, result)
|
||||
return
|
||||
if path == "/api/training/jobs":
|
||||
self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload()))
|
||||
return
|
||||
@@ -859,6 +1113,11 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument(
|
||||
"--trainer-python", default=sys.executable, help="已安装 mjlab/torch 的 Python 解释器"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mobile-python",
|
||||
default=None,
|
||||
help="可选独立移动操作 Python(MuJoCo 3.11.0 + SB3),避免更改 Go2 环境",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tuning-data-root",
|
||||
type=Path,
|
||||
@@ -902,6 +1161,7 @@ def main() -> None:
|
||||
args.trainer_root,
|
||||
args.trainer_python,
|
||||
tuple(args.tasks or DEFAULT_TASKS),
|
||||
mobile_python=args.mobile_python,
|
||||
lease=lease,
|
||||
sources=sources,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user