feat(training): release V0.8 自调参 Agent
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-02 13:49:34 +08:00
parent cffac29a03
commit deead17a9a
47 changed files with 4986 additions and 96 deletions
+164 -24
View File
@@ -23,9 +23,12 @@ from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Any
from urllib.parse import unquote, urlsplit
from urllib.parse import parse_qs, unquote, urlsplit
VERSION = "0.3.0"
from tuning.manager import TuningError, TuningManager
from tuning.process import GpuLease, ResourceBusyError
VERSION = "0.4.0"
# 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。
DEFAULT_TASKS = ("Unitree-Go2-Flat",)
ACTIVE_STATES = {"queued", "running"}
@@ -65,6 +68,7 @@ class TrainingConfig:
device: str
gpu_ids: list[int]
wandb_mode: str
reward_config: dict[str, Any] | None = None
@dataclass
@@ -110,6 +114,7 @@ class TrainingManager:
python: str,
tasks: tuple[str, ...],
check_environment: bool = True,
lease: GpuLease | None = None,
):
self.trainer_root = trainer_root.expanduser().resolve()
self.python = str(Path(python).expanduser()) if os.sep in python else python
@@ -118,6 +123,8 @@ class TrainingManager:
self.lock = threading.RLock()
self.check_environment = check_environment
self._environment_error: str | None | bool = False
self.lease = lease or GpuLease()
self.preset_resolver: Any = None
def readiness_error(self) -> str | None:
if not self.trainer_root.is_dir():
@@ -204,6 +211,17 @@ class TrainingManager:
wandb_mode = payload.get("wandbMode", "offline")
if wandb_mode not in ("offline", "disabled", "online"):
raise ApiError(HTTPStatus.BAD_REQUEST, "wandbMode 必须是 offline、disabled 或 online")
preset_id = payload.get("rewardPresetId")
reward_config = None
if preset_id is not None:
if not isinstance(preset_id, str) or not re.fullmatch(r"[0-9a-f]{32}", preset_id):
raise ApiError(HTTPStatus.BAD_REQUEST, "rewardPresetId 格式无效")
if self.preset_resolver is None:
raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 服务未就绪")
try:
reward_config = self.preset_resolver(preset_id)
except KeyError as error:
raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 不存在") from error
return TrainingConfig(
task_id=task_id,
num_envs=integer("numEnvs", 1, 16384),
@@ -213,6 +231,7 @@ class TrainingManager:
device=device,
gpu_ids=raw_gpu_ids,
wandb_mode=wandb_mode,
reward_config=reward_config,
)
def start(self, payload: Any) -> dict[str, Any]:
@@ -232,10 +251,20 @@ class TrainingManager:
raise ApiError(HTTPStatus.CONFLICT, "训练任务历史已满,请稍后重试")
del self.jobs[completed]
job = TrainingJob(id=uuid.uuid4().hex, config=config)
owner = f"training:{job.id}"
try:
self.lease.acquire(owner)
except ResourceBusyError as error:
raise ApiError(HTTPStatus.CONFLICT, str(error)) from error
self.jobs[job.id] = job
threading.Thread(
target=self._run, args=(job,), name=f"training-{job.id[:8]}", daemon=True
).start()
try:
threading.Thread(
target=self._run, args=(job,), name=f"training-{job.id[:8]}", daemon=True
).start()
except Exception:
self.jobs.pop(job.id, None)
self.lease.release(owner)
raise
return job.public()
def get(self, job_id: str) -> dict[str, Any]:
@@ -305,6 +334,13 @@ class TrainingManager:
f"--agent.seed={config.seed}",
f"--agent.run-name={config.run_name}",
]
if config.reward_config is not None:
command.extend(
(
"--reward-config-json",
json.dumps(config.reward_config, ensure_ascii=False, separators=(",", ":")),
)
)
if config.device == "cpu":
command.extend(("--gpu-ids", "None"))
else:
@@ -408,13 +444,16 @@ class TrainingManager:
job.state = "cancelled" if job.cancel_requested else "failed"
job.message = f"启动训练失败:{error}"
job.logs.append(job.message)
finally:
self.lease.release(f"training:{job.id}")
class TrainingRequestHandler(BaseHTTPRequestHandler):
manager: TrainingManager
tuning_manager: TuningManager
allowed_origins: tuple[str, ...] = ()
access_token = ""
server_version = "MuJoCoLocalTraining/0.3"
server_version = "MuJoCoLocalTraining/0.4"
def log_message(self, format: str, *args: Any) -> None:
sys.stderr.write(f"[{self.log_date_time_string()}] {format % args}\n")
@@ -455,6 +494,12 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
def _error(self, error: Exception) -> None:
if isinstance(error, ApiError):
self._json(error.status, {"error": str(error)})
elif isinstance(error, KeyError):
self._json(HTTPStatus.NOT_FOUND, {"error": "调参 session、trial 或 proposal 不存在"})
elif isinstance(error, ResourceBusyError):
self._json(HTTPStatus.CONFLICT, {"error": str(error)})
elif isinstance(error, TuningError):
self._json(HTTPStatus.BAD_REQUEST, {"error": str(error)})
else:
self._json(
HTTPStatus.INTERNAL_SERVER_ERROR, {"error": f"本地训练服务内部错误:{error}"}
@@ -490,6 +535,18 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
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 _send_file(self, file_path: Path, filename: str) -> None:
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", f'attachment; filename="{filename}"')
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)
def do_OPTIONS(self) -> None:
try:
self._ensure_origin()
@@ -505,25 +562,57 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
def do_GET(self) -> None:
try:
self._ensure_request()
path = urlsplit(self.path).path
parsed = urlsplit(self.path)
path = parsed.path
if path == "/api/training/health":
self._json(HTTPStatus.OK, self.manager.health())
health = self.manager.health()
health["tuning"] = self.tuning_manager.capability()
health["resourceOwner"] = self.manager.lease.public()
self._json(HTTPStatus.OK, health)
return
if path == "/api/tuning/capabilities":
self._json(HTTPStatus.OK, self.tuning_manager.capability())
return
if path == "/api/tuning/sessions":
self._json(HTTPStatus.OK, {"sessions": self.tuning_manager.list()})
return
if path == "/api/tuning/presets":
self._json(HTTPStatus.OK, {"presets": self.tuning_manager.storage.list_presets()})
return
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})", path)
if match:
self._json(HTTPStatus.OK, self.tuning_manager.detail(match.group(1)))
return
match = re.fullmatch(
r"/api/tuning/sessions/([0-9a-f]{32})/trials/([0-9a-f]{32})/metrics", path
)
if match:
query = parse_qs(parsed.query)
tags = [tag for value in query.get("tags", []) for tag in value.split(",") if tag]
try:
max_points = int(query.get("maxPoints", ["1000"])[0])
except ValueError as error:
raise TuningError("maxPoints 必须是整数") from error
if not 10 <= max_points <= 5000:
raise TuningError("maxPoints 必须在 10–5000 之间")
self._json(
HTTPStatus.OK,
self.tuning_manager.metrics(
match.group(1), match.group(2), tags or None, max_points
),
)
return
match = re.fullmatch(
r"/api/tuning/sessions/([0-9a-f]{32})/artifacts/best/policy\.onnx", path
)
if match:
self._send_file(self.tuning_manager.best_artifact(match.group(1)), "policy.onnx")
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)
self._send_file(self.manager.artifact(job_id), "policy.onnx")
else:
self._json(HTTPStatus.OK, self.manager.get(job_id))
except Exception as error:
@@ -532,16 +621,52 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
try:
self._ensure_request()
if urlsplit(self.path).path != "/api/training/jobs":
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload()))
path = urlsplit(self.path).path
if path == "/api/training/jobs":
self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload()))
return
if path == "/api/tuning/agent/test":
self._json(HTTPStatus.OK, self.tuning_manager.test_agent())
return
if path == "/api/tuning/sessions":
self._json(HTTPStatus.ACCEPTED, self.tuning_manager.create(self._payload()))
return
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})/(pause|resume)", path)
if match:
action = (
self.tuning_manager.pause
if match.group(2) == "pause"
else self.tuning_manager.resume
)
self._json(HTTPStatus.ACCEPTED, action(match.group(1)))
return
match = re.fullmatch(
r"/api/tuning/sessions/([0-9a-f]{32})/proposals/([0-9a-f]{32})/(approve|reject)",
path,
)
if match:
action = (
self.tuning_manager.approve
if match.group(3) == "approve"
else self.tuning_manager.reject
)
self._json(
HTTPStatus.ACCEPTED, action(match.group(1), match.group(2), self._payload())
)
return
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
except Exception as error:
self._error(error)
def do_DELETE(self) -> None:
try:
self._ensure_request()
job_id, artifact = self._route(urlsplit(self.path).path)
path = urlsplit(self.path).path
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})", path)
if match:
self._json(HTTPStatus.ACCEPTED, self.tuning_manager.cancel(match.group(1)))
return
job_id, artifact = self._route(path)
if not job_id or artifact:
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
self._json(HTTPStatus.ACCEPTED, self.manager.cancel(job_id))
@@ -574,6 +699,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument(
"--trainer-python", default=sys.executable, help="已安装 mjlab/torch 的 Python 解释器"
)
parser.add_argument(
"--tuning-data-root",
type=Path,
default=None,
help="调参 SQLite 与 trial 产物目录;默认位于训练工程 logs/auto_tuning",
)
parser.add_argument(
"--task", action="append", dest="tasks", help="允许前端启动的任务 ID;可重复"
)
@@ -593,10 +724,18 @@ def main() -> None:
token = args.token or secrets.token_urlsafe(24)
if len(token) < 16:
raise SystemExit("训练服务访问令牌至少需要 16 个字符")
lease = GpuLease()
manager = TrainingManager(
args.trainer_root, args.trainer_python, tuple(args.tasks or DEFAULT_TASKS)
args.trainer_root,
args.trainer_python,
tuple(args.tasks or DEFAULT_TASKS),
lease=lease,
)
tuning_root = args.tuning_data_root or (Path(args.trainer_root) / "logs" / "auto_tuning")
tuning_manager = TuningManager(args.trainer_root, args.trainer_python, tuning_root, lease)
manager.preset_resolver = tuning_manager.preset_config
TrainingRequestHandler.manager = manager
TrainingRequestHandler.tuning_manager = tuning_manager
TrainingRequestHandler.allowed_origins = tuple(args.allow_origin)
TrainingRequestHandler.access_token = token
server = ThreadingHTTPServer((args.host, args.port), TrainingRequestHandler)
@@ -612,6 +751,7 @@ def main() -> None:
except KeyboardInterrupt:
print("\n正在停止本地训练服务…")
finally:
tuning_manager.shutdown()
manager.shutdown()
server.server_close()