feat(training): release V0.8 自调参 Agent
This commit is contained in:
+164
-24
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user