feat(training): release V0.8 自调参 Agent
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
# 本地强化学习训练服务
|
||||
|
||||
该服务把 Web 平台发出的受限训练请求转换为本机训练子进程,并提供状态轮询、停止任务和 `policy.onnx` 下载接口。服务只绑定 `127.0.0.1`,不执行前端传入的任意命令。
|
||||
该服务把 Web 平台发出的受限训练请求转换为本机训练子进程,并提供状态轮询、停止任务和 `policy.onnx` 下载接口。它还提供 `Unitree-Go2-Flat` 奖励函数自调参:DeepSeek Agent 根据训练曲线与固定评估指标提出受限参数 patch,系统支持全自动或逐轮审批、successive-halving、TensorBoard scalar 查询、最佳 preset 与 ONNX 导出。服务只绑定 `127.0.0.1`,不执行前端传入的任意命令或 Agent 生成的代码。
|
||||
|
||||
仓库已在 [`rl/`](rl/) 内置 `Unitree-Go2-Flat` 所需的 PPO 训练代码、Go2 模型资产和 ONNX 导出逻辑,不再要求另外克隆 `unitree_rl_mjlab`。`mjlab`、PyTorch 等大型运行依赖仍需安装在本机训练环境中。
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
python -m pip install -r training_server/rl/requirements.txt
|
||||
python -m pip install -r training_server/requirements.txt
|
||||
```
|
||||
|
||||
## 启动
|
||||
@@ -29,6 +30,17 @@ MUJOCO_TRAINING_TOKEN='至少十六个字符的随机令牌' npm run training-se
|
||||
--trainer-python /path/to/training-env/bin/python
|
||||
```
|
||||
|
||||
启用云端自调参 Agent 时,在**服务端环境变量**中配置 DeepSeek;不要把 key 填到浏览器、URL 或命令行参数:
|
||||
|
||||
```bash
|
||||
export DEEPSEEK_API_KEY='你的 DeepSeek API key'
|
||||
export MUJOCO_TUNING_AGENT_BASE_URL='https://api.deepseek.com' # 可省略
|
||||
export MUJOCO_TUNING_AGENT_MODEL='deepseek-v4-flash' # 可省略
|
||||
npm run training-server -- --trainer-python "$PWD/.venv/bin/python"
|
||||
```
|
||||
|
||||
普通训练不要求 DeepSeek key。未配置时健康接口会把 tuning 标记为不可用;只有创建 session 时明确勾选 fallback,才允许 Agent 失败后使用 Optuna 候选,不会静默降级。
|
||||
|
||||
默认训练工程是仓库内的 `training_server/rl`。如需使用包含其他已注册任务的外部训练工程,仍可通过 `--trainer-root /path/to/trainer` 或 `UNITREE_RL_MJLAB_ROOT` 覆盖。默认端口是 `8765`。如果前端不是从 `localhost` 或 `127.0.0.1` 提供,可显式添加来源:
|
||||
|
||||
```bash
|
||||
@@ -37,7 +49,17 @@ python training_server/server.py \
|
||||
--allow-origin http://192.168.1.10:5173
|
||||
```
|
||||
|
||||
服务一次只运行一个训练任务,最多保留 20 个任务的内存状态,每个任务最多保留 200 行最近日志。停止服务或在界面点击“停止训练”会同步终止整个训练进程组。训练请求的 W&B 模式默认为 `offline`,保留本地指标但不登录;也可以在界面选择完全禁用或在线模式。所有 API 请求都必须携带启动时生成的 Bearer Token。
|
||||
普通训练与自调参共享同一个计算资源锁,任何时刻只允许一个训练/评估子进程占用 GPU。普通任务最多保留 20 个内存状态和每个任务 200 行最近日志;调参 session、trial、proposal、审计和 scalar 写入 SQLite/WAL,默认保存在 `training_server/rl/logs/auto_tuning/`。服务重启后等待审批/暂停状态可恢复,正在训练或评估的 trial 标记为 interrupted,只能从已完整保存的 checkpoint 显式恢复。停止服务或点击“停止”会终止整个进程组。
|
||||
|
||||
普通训练的 W&B 模式默认为 `offline`;调参 trial 强制使用本地 TensorBoard writer,DeepSeek 只接收最多 12 个 trial 的脱敏数值摘要和降采样曲线,不接收源代码、机器人资产、checkpoint、服务 token 或本地路径。所有 HTTP API 请求都必须携带启动时生成的 Bearer Token。
|
||||
|
||||
## 自调参流程
|
||||
|
||||
在主工作台连接训练服务后,点击“打开自调参 Agent 工作台”。默认预算为 12 个唯一配置:所有配置先训练 300 iterations,前 4 名续训到 900,前 2 名续训到 2000;默认使用 GPU 0 和 4096 个并行环境。首次使用建议先降低为 256–512 environments 做 smoke test。
|
||||
|
||||
固定评估使用站立、前进/侧移、转向和组合命令以及 3 个固定 seed。最终分数不直接使用可被权重放大的总 reward,而由速度跟踪 35%、动作平滑 20%、姿态稳定 15%、减少跌倒 15%、足端滑移 10%、能耗 5% 的权重无关指标组成。跌倒率高于基线 2% 或速度误差恶化超过 5% 的 trial 不晋级。逐轮审批模式会自动运行基线,之后每条 Agent 建议都等待批准、修改后批准或拒绝反馈。
|
||||
|
||||
最佳结果保存为不可变 preset,可在普通训练面板的“奖励配置”中选择,也可导出 JSON;不会覆盖仓库里的 Python 默认奖励配置。
|
||||
|
||||
## 接口
|
||||
|
||||
@@ -45,14 +67,22 @@ python training_server/server.py \
|
||||
- `POST /api/training/jobs`:发起训练;
|
||||
- `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志;
|
||||
- `DELETE /api/training/jobs/{id}`:停止训练;
|
||||
- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略。
|
||||
- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略;
|
||||
- `GET /api/tuning/capabilities`、`POST /api/tuning/agent/test`:检查/测试 Agent;
|
||||
- `GET|POST /api/tuning/sessions`、`GET|DELETE /api/tuning/sessions/{id}`:列出、创建、查询、停止 session;
|
||||
- `POST /api/tuning/sessions/{id}/pause|resume`:暂停后续调度或恢复;
|
||||
- `POST /api/tuning/sessions/{id}/proposals/{proposalId}/approve|reject`:审批、修改或拒绝建议;
|
||||
- `GET /api/tuning/sessions/{id}/trials/{trialId}/metrics`:查询降采样 scalar;
|
||||
- `GET /api/tuning/sessions/{id}/artifacts/best/policy.onnx`:下载最佳策略;
|
||||
- `GET /api/tuning/presets`:列出可供普通训练复用的最佳奖励 preset。
|
||||
|
||||
任务保存在服务内存中,服务重启后历史任务状态会丢失;训练日志、checkpoint 和 ONNX 产物保留在 `training_server/rl/logs/rsl_rl/`,使用外部训练工程时则保留在对应工程中。
|
||||
普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态及产物持久化在 `logs/auto_tuning/`。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、每轮最多 4 项及 `0.5×–2×` 变化率校验。
|
||||
|
||||
## 测试
|
||||
|
||||
```bash
|
||||
python3 -m pip install -r requirements-dev.txt
|
||||
source .venv/bin/activate
|
||||
python -m pip install -r requirements-dev.txt
|
||||
npm run lint:python
|
||||
npm run test:training-server
|
||||
```
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# 本地自调参服务(Python 3.12)
|
||||
pydantic-ai-slim[openai]==2.37.0
|
||||
httpx2[socks]==2.12.0
|
||||
optuna==4.9.0
|
||||
tensorboard==2.21.0
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Deterministic, headless evaluation for Unitree Go2 velocity policies."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from pathlib import Path
|
||||
from statistics import fmean, pstdev
|
||||
from typing import Literal
|
||||
|
||||
TRAINER_ROOT = Path(__file__).resolve().parents[1]
|
||||
SERVICE_ROOT = TRAINER_ROOT.parent
|
||||
for source_root in (TRAINER_ROOT, SERVICE_ROOT):
|
||||
if str(source_root) not in sys.path:
|
||||
sys.path.insert(0, str(source_root))
|
||||
|
||||
import torch
|
||||
import tyro
|
||||
import warp as wp
|
||||
|
||||
if not hasattr(wp, "context"):
|
||||
from warp._src import context as warp_context
|
||||
|
||||
wp.context = warp_context # type: ignore[attr-defined]
|
||||
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from mjlab.rl import MjlabOnPolicyRunner, RslRlVecEnvWrapper
|
||||
from mjlab.tasks.registry import list_tasks, load_env_cfg, load_rl_cfg, load_runner_cls
|
||||
from mjlab.tasks.velocity.mdp import UniformVelocityCommandCfg
|
||||
from mjlab.utils.torch import configure_torch_backends
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from tuning.schema import apply_reward_configuration, validate_configuration
|
||||
|
||||
SCENARIOS = (
|
||||
(0.0, 0.0, 0.0),
|
||||
(0.5, 0.0, 0.0),
|
||||
(1.0, 0.0, 0.0),
|
||||
(1.5, 0.0, 0.0),
|
||||
(0.0, 0.5, 0.0),
|
||||
(0.0, -0.5, 0.0),
|
||||
(0.0, 0.0, 0.5),
|
||||
(0.0, 0.0, -0.5),
|
||||
(0.8, 0.25, 0.35),
|
||||
)
|
||||
METRIC_NAMES = (
|
||||
"linear_velocity_rmse",
|
||||
"angular_velocity_rmse",
|
||||
"mean_action_acc",
|
||||
"orientation_error",
|
||||
"slip_velocity",
|
||||
"mechanical_power",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EvaluateConfig:
|
||||
checkpoint: str
|
||||
output: str
|
||||
reward_config: str | None = None
|
||||
num_envs: int = 256
|
||||
steps_per_seed: int = 1000
|
||||
seeds: tuple[int, ...] = field(default_factory=lambda: (101, 202, 303))
|
||||
device: str | None = None
|
||||
gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0])
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _fixed_commands(command: torch.Tensor) -> torch.Tensor:
|
||||
values = torch.as_tensor(SCENARIOS, device=command.device, dtype=command.dtype)
|
||||
indexes = torch.arange(command.shape[0], device=command.device) % values.shape[0]
|
||||
command[:] = values[indexes]
|
||||
return command
|
||||
|
||||
|
||||
def _evaluate_seed(task_id: str, cfg: EvaluateConfig, seed: int) -> dict[str, float]:
|
||||
torch.manual_seed(seed)
|
||||
env_cfg = load_env_cfg(task_id, play=False)
|
||||
agent_cfg = load_rl_cfg(task_id)
|
||||
env_cfg.seed = seed
|
||||
env_cfg.scene.num_envs = cfg.num_envs
|
||||
env_cfg.curriculum = {}
|
||||
env_cfg.observations["actor"].enable_corruption = False
|
||||
env_cfg.events.pop("push_robot", None)
|
||||
twist_cfg = env_cfg.commands["twist"]
|
||||
assert isinstance(twist_cfg, UniformVelocityCommandCfg)
|
||||
twist_cfg.heading_command = False
|
||||
twist_cfg.ranges.heading = None
|
||||
twist_cfg.rel_heading_envs = 0.0
|
||||
twist_cfg.rel_standing_envs = 0.0
|
||||
twist_cfg.resampling_time_range = (1.0e9, 1.0e9)
|
||||
if cfg.reward_config:
|
||||
reward_path = Path(cfg.reward_config).expanduser().resolve(strict=True)
|
||||
with reward_path.open(encoding="utf-8") as stream:
|
||||
apply_reward_configuration(env_cfg, validate_configuration(json.load(stream)))
|
||||
|
||||
env = ManagerBasedRlEnv(cfg=env_cfg, device=cfg.device or "cuda:0")
|
||||
wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
try:
|
||||
runner_cls = load_runner_cls(task_id) or MjlabOnPolicyRunner
|
||||
runner = runner_cls(wrapped, asdict(agent_cfg), log_dir=None, device=wrapped.device)
|
||||
runner.load(
|
||||
str(Path(cfg.checkpoint).expanduser().resolve(strict=True)),
|
||||
load_cfg={"actor": True},
|
||||
strict=True,
|
||||
map_location=str(wrapped.device),
|
||||
)
|
||||
policy = runner.get_inference_policy(device=str(wrapped.device))
|
||||
twist = wrapped.unwrapped.command_manager.get_term("twist")
|
||||
_fixed_commands(twist.command)
|
||||
obs = wrapped.get_observations()
|
||||
|
||||
sums = {name: 0.0 for name in METRIC_NAMES}
|
||||
samples = 0
|
||||
terminations = 0
|
||||
completions = 0
|
||||
with torch.inference_mode():
|
||||
for _ in range(cfg.steps_per_seed):
|
||||
_fixed_commands(twist.command)
|
||||
obs = wrapped.get_observations()
|
||||
actions = policy(obs)
|
||||
obs, _rewards, _dones, _extras = wrapped.step(actions)
|
||||
manager = wrapped.unwrapped.metrics_manager
|
||||
for index, name in enumerate(manager.active_terms):
|
||||
if name in sums:
|
||||
sums[name] += float(torch.sum(manager._step_values[:, index]).item())
|
||||
samples += wrapped.num_envs
|
||||
terminated = wrapped.unwrapped.reset_terminated
|
||||
timed_out = wrapped.unwrapped.reset_time_outs
|
||||
terminations += int(torch.count_nonzero(terminated).item())
|
||||
completions += int(torch.count_nonzero(terminated | timed_out).item())
|
||||
result = {name: sums[name] / max(samples, 1) for name in METRIC_NAMES}
|
||||
result["fall_rate"] = terminations / max(completions, wrapped.num_envs)
|
||||
return result
|
||||
finally:
|
||||
wrapped.close()
|
||||
|
||||
|
||||
def run_evaluation(task_id: str, cfg: EvaluateConfig) -> dict:
|
||||
configure_torch_backends()
|
||||
selected = cfg.gpu_ids
|
||||
if selected is None:
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
device = "cpu"
|
||||
else:
|
||||
if selected == "all":
|
||||
selected = [0]
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ",".join(map(str, selected))
|
||||
device = cfg.device or "cuda:0"
|
||||
os.environ["MUJOCO_GL"] = "egl"
|
||||
cfg = EvaluateConfig(**{**asdict(cfg), "device": device})
|
||||
|
||||
checkpoint = Path(cfg.checkpoint).expanduser().resolve(strict=True)
|
||||
output = Path(cfg.output).expanduser().resolve()
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
per_seed = [_evaluate_seed(task_id, cfg, seed) for seed in cfg.seeds]
|
||||
metrics = {
|
||||
key: fmean(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
deviations = {
|
||||
key: pstdev(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
result = {
|
||||
"protocolVersion": 1,
|
||||
"taskId": task_id,
|
||||
"checkpoint": checkpoint.name,
|
||||
"checkpointSha256": _sha256(checkpoint),
|
||||
"seeds": list(cfg.seeds),
|
||||
"numEnvs": cfg.num_envs,
|
||||
"stepsPerSeed": cfg.steps_per_seed,
|
||||
"scenarios": [list(value) for value in SCENARIOS],
|
||||
"metrics": metrics,
|
||||
"metricStd": deviations,
|
||||
"seedMetrics": [
|
||||
{"seed": seed, "metrics": values}
|
||||
for seed, values in zip(cfg.seeds, per_seed, strict=True)
|
||||
],
|
||||
}
|
||||
output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
writer = SummaryWriter(log_dir=str(output.parent / "evaluation-events"))
|
||||
try:
|
||||
for name, value in metrics.items():
|
||||
writer.add_scalar(f"Evaluation/{name}", value, 0)
|
||||
finally:
|
||||
writer.close()
|
||||
print("MUJOCO_EVALUATION " + json.dumps({"output": str(output), "metrics": metrics}))
|
||||
return result
|
||||
|
||||
|
||||
def main() -> None:
|
||||
import mjlab.tasks # noqa: F401
|
||||
import src.tasks # noqa: F401
|
||||
|
||||
chosen_task, remaining = tyro.cli(
|
||||
tyro.extras.literal_type_from_choices(list_tasks()),
|
||||
add_help=False,
|
||||
return_unknown_args=True,
|
||||
config=mjlab.TYRO_FLAGS,
|
||||
)
|
||||
args = tyro.cli(
|
||||
EvaluateConfig,
|
||||
args=remaining,
|
||||
prog=sys.argv[0] + f" {chosen_task}",
|
||||
config=mjlab.TYRO_FLAGS,
|
||||
)
|
||||
run_evaluation(chosen_task, args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Script to train RL agent with RSL-RL."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
@@ -10,8 +11,10 @@ from typing import Literal, cast
|
||||
|
||||
# 训练器作为仓库内置子集直接从 scripts/ 启动,不要求额外执行 pip install -e。
|
||||
TRAINER_ROOT = Path(__file__).resolve().parents[1]
|
||||
if str(TRAINER_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(TRAINER_ROOT))
|
||||
SERVICE_ROOT = TRAINER_ROOT.parent
|
||||
for source_root in (TRAINER_ROOT, SERVICE_ROOT):
|
||||
if str(source_root) not in sys.path:
|
||||
sys.path.insert(0, str(source_root))
|
||||
|
||||
import tyro
|
||||
import warp as wp
|
||||
@@ -32,6 +35,8 @@ from mjlab.utils.os import dump_yaml, get_checkpoint_path
|
||||
from mjlab.utils.torch import configure_torch_backends
|
||||
from mjlab.utils.wrappers import VideoRecorder
|
||||
|
||||
from tuning.schema import apply_reward_configuration, validate_configuration
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TrainConfig:
|
||||
@@ -44,6 +49,10 @@ class TrainConfig:
|
||||
enable_nan_guard: bool = False
|
||||
torchrunx_log_dir: str | None = None
|
||||
gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0])
|
||||
output_dir: str | None = None
|
||||
resume_checkpoint: str | None = None
|
||||
reward_config: str | None = None
|
||||
reward_config_json: str | None = None
|
||||
|
||||
@staticmethod
|
||||
def from_task(task_id: str) -> "TrainConfig":
|
||||
@@ -52,7 +61,27 @@ class TrainConfig:
|
||||
return TrainConfig(env=env_cfg, agent=agent_cfg)
|
||||
|
||||
|
||||
def _load_reward_config(path: str | None, inline: str | None) -> dict | None:
|
||||
if path is not None and inline is not None:
|
||||
raise ValueError("Use only one of reward_config and reward_config_json")
|
||||
if inline is not None:
|
||||
if len(inline.encode("utf-8")) > 64 * 1024:
|
||||
raise ValueError("Reward configuration is larger than 64 KiB")
|
||||
return validate_configuration(json.loads(inline))
|
||||
if path is None:
|
||||
return None
|
||||
source = Path(path).expanduser().resolve(strict=True)
|
||||
if source.stat().st_size > 64 * 1024:
|
||||
raise ValueError("Reward configuration is larger than 64 KiB")
|
||||
with source.open(encoding="utf-8") as stream:
|
||||
return validate_configuration(json.load(stream))
|
||||
|
||||
|
||||
def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
reward_config = _load_reward_config(cfg.reward_config, cfg.reward_config_json)
|
||||
if reward_config is not None:
|
||||
apply_reward_configuration(cfg.env, reward_config)
|
||||
|
||||
cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "")
|
||||
if cuda_visible == "":
|
||||
device = "cpu"
|
||||
@@ -109,11 +138,14 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
log_root_path = log_dir.parent # Go up from specific run dir to experiment dir.
|
||||
|
||||
resume_path: Path | None = None
|
||||
if cfg.agent.resume:
|
||||
# Load checkpoint from local filesystem.
|
||||
resume_path = get_checkpoint_path(
|
||||
log_root_path, cfg.agent.load_run, cfg.agent.load_checkpoint
|
||||
)
|
||||
explicit_resume = cfg.resume_checkpoint is not None
|
||||
if explicit_resume:
|
||||
resume_path = Path(cfg.resume_checkpoint).expanduser().resolve(strict=True)
|
||||
elif cfg.agent.resume:
|
||||
# Load checkpoint from local filesystem.
|
||||
resume_path = get_checkpoint_path(
|
||||
log_root_path, cfg.agent.load_run, cfg.agent.load_checkpoint
|
||||
)
|
||||
|
||||
# Only record videos on rank 0 to avoid multiple workers writing to the same files.
|
||||
if cfg.video and rank == 0:
|
||||
@@ -141,16 +173,32 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
runner.add_git_repo_to_log(__file__)
|
||||
if resume_path is not None:
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
runner.load(str(resume_path))
|
||||
runner.load(str(resume_path), map_location=device)
|
||||
if explicit_resume:
|
||||
# RSL-RL stores the last completed zero-based iteration and otherwise
|
||||
# repeats it after load. Explicit tuning promotion uses an absolute target.
|
||||
runner.current_learning_iteration += 1
|
||||
|
||||
# Only write config files from rank 0 to avoid race conditions.
|
||||
if rank == 0:
|
||||
dump_yaml(log_dir / "params" / "env.yaml", env_cfg)
|
||||
dump_yaml(log_dir / "params" / "agent.yaml", agent_cfg)
|
||||
if reward_config is not None:
|
||||
reward_snapshot = log_dir / "params" / "reward_config.json"
|
||||
reward_snapshot.parent.mkdir(parents=True, exist_ok=True)
|
||||
reward_snapshot.write_text(
|
||||
json.dumps(reward_config, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||
)
|
||||
|
||||
runner.learn(
|
||||
num_learning_iterations=cfg.agent.max_iterations, init_at_random_ep_len=True
|
||||
iterations = cfg.agent.max_iterations
|
||||
if explicit_resume:
|
||||
iterations = max(0, cfg.agent.max_iterations - runner.current_learning_iteration)
|
||||
print(
|
||||
f"[INFO] Learning target: current={runner.current_learning_iteration}, "
|
||||
f"additional={iterations}, target={cfg.agent.max_iterations}",
|
||||
flush=True,
|
||||
)
|
||||
runner.learn(num_learning_iterations=iterations, init_at_random_ep_len=True)
|
||||
|
||||
env.close()
|
||||
|
||||
@@ -159,12 +207,15 @@ def launch_training(task_id: str, args: TrainConfig | None = None):
|
||||
args = args or TrainConfig.from_task(task_id)
|
||||
|
||||
# Create log directory once before launching workers.
|
||||
log_root_path = Path("logs") / "rsl_rl" / args.agent.experiment_name
|
||||
log_root_path.resolve()
|
||||
log_dir_name = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
if args.agent.run_name:
|
||||
log_dir_name += f"_{args.agent.run_name}"
|
||||
log_dir = log_root_path / log_dir_name
|
||||
if args.output_dir:
|
||||
log_dir = Path(args.output_dir).expanduser().resolve()
|
||||
else:
|
||||
log_root_path = (Path("logs") / "rsl_rl" / args.agent.experiment_name).resolve()
|
||||
log_dir_name = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
if args.agent.run_name:
|
||||
log_dir_name += f"_{args.agent.run_name}"
|
||||
log_dir = log_root_path / log_dir_name
|
||||
log_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Select GPUs based on CUDA_VISIBLE_DEVICES and user specification.
|
||||
selected_gpus, num_gpus = select_gpus(args.gpu_ids)
|
||||
|
||||
@@ -106,6 +106,7 @@ def unitree_go2_rough_env_cfg(
|
||||
cfg.rewards["body_ang_vel"].params["asset_cfg"].body_names = ("base_link",)
|
||||
cfg.rewards["foot_clearance"].params["asset_cfg"].site_names = site_names
|
||||
cfg.rewards["foot_slip"].params["asset_cfg"].site_names = site_names
|
||||
cfg.metrics["slip_velocity"].params["asset_cfg"].site_names = site_names
|
||||
|
||||
cfg.terminations["illegal_contact"] = TerminationTermCfg(
|
||||
func=mdp.illegal_contact,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from mjlab.envs.mdp import * # noqa: F401, F403
|
||||
|
||||
from .curriculums import * # noqa: F403
|
||||
from .metrics import * # noqa: F403
|
||||
from .observations import * # noqa: F403
|
||||
from .rewards import * # noqa: F403
|
||||
from .terminations import * # noqa: F403
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Reward-weight-independent quality metrics for velocity tasks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from mjlab.entity import Entity
|
||||
from mjlab.managers.scene_entity_config import SceneEntityCfg
|
||||
from mjlab.sensor import ContactSensor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
|
||||
_DEFAULT_ASSET_CFG = SceneEntityCfg("robot")
|
||||
|
||||
|
||||
def linear_velocity_rmse(
|
||||
env: ManagerBasedRlEnv,
|
||||
command_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Per-step commanded-vs-actual base linear velocity RMSE in body frame."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
command = env.command_manager.get_command(command_name)
|
||||
assert command is not None
|
||||
actual = asset.data.root_link_lin_vel_b
|
||||
error = torch.cat((command[:, :2] - actual[:, :2], -actual[:, 2:3]), dim=1)
|
||||
return torch.sqrt(torch.mean(torch.square(error), dim=1))
|
||||
|
||||
|
||||
def angular_velocity_rmse(
|
||||
env: ManagerBasedRlEnv,
|
||||
command_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Per-step commanded-vs-actual base angular velocity RMSE in body frame."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
command = env.command_manager.get_command(command_name)
|
||||
assert command is not None
|
||||
actual = asset.data.root_link_ang_vel_b
|
||||
desired = torch.zeros_like(actual)
|
||||
desired[:, 2] = command[:, 2]
|
||||
return torch.sqrt(torch.mean(torch.square(desired - actual), dim=1))
|
||||
|
||||
|
||||
def orientation_error(
|
||||
env: ManagerBasedRlEnv,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Magnitude of projected gravity in the base x/y plane; zero is upright."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
return torch.linalg.vector_norm(asset.data.projected_gravity_b[:, :2], dim=1)
|
||||
|
||||
|
||||
def fall_indicator(env: ManagerBasedRlEnv) -> torch.Tensor:
|
||||
"""One on non-timeout terminal steps, otherwise zero."""
|
||||
return env.termination_manager.terminated.float()
|
||||
|
||||
|
||||
def slip_velocity(
|
||||
env: ManagerBasedRlEnv,
|
||||
sensor_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Mean x/y velocity of feet currently touching the ground."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
sensor: ContactSensor = env.scene[sensor_name]
|
||||
assert sensor.data.found is not None
|
||||
contact = (sensor.data.found > 0).float()
|
||||
speed = torch.linalg.vector_norm(
|
||||
asset.data.site_lin_vel_w[:, asset_cfg.site_ids, :2], dim=-1
|
||||
)
|
||||
count = torch.clamp(torch.sum(contact, dim=1), min=1.0)
|
||||
return torch.sum(speed * contact, dim=1) / count
|
||||
|
||||
|
||||
def mechanical_power(
|
||||
env: ManagerBasedRlEnv,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Positive actuator mechanical power in watts (regeneration is ignored)."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
torque = asset.data.actuator_force[:, asset_cfg.actuator_ids]
|
||||
velocity = asset.data.joint_vel[:, asset_cfg.joint_ids]
|
||||
count = min(torque.shape[1], velocity.shape[1])
|
||||
return torch.sum(torch.clamp(torque[:, :count] * velocity[:, :count], min=0.0), dim=1)
|
||||
@@ -140,8 +140,27 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg:
|
||||
##
|
||||
|
||||
metrics = {
|
||||
"mean_action_acc": MetricsTermCfg(
|
||||
func=mdp.mean_action_acc,
|
||||
"linear_velocity_rmse": MetricsTermCfg(
|
||||
func=mdp.linear_velocity_rmse,
|
||||
params={"command_name": "twist"},
|
||||
),
|
||||
"angular_velocity_rmse": MetricsTermCfg(
|
||||
func=mdp.angular_velocity_rmse,
|
||||
params={"command_name": "twist"},
|
||||
),
|
||||
"mean_action_acc": MetricsTermCfg(func=mdp.mean_action_acc),
|
||||
"orientation_error": MetricsTermCfg(func=mdp.orientation_error),
|
||||
"fall_indicator": MetricsTermCfg(func=mdp.fall_indicator),
|
||||
"slip_velocity": MetricsTermCfg(
|
||||
func=mdp.slip_velocity,
|
||||
params={
|
||||
"sensor_name": "feet_ground_contact",
|
||||
"asset_cfg": SceneEntityCfg("robot", site_names=()), # Set per-robot.
|
||||
},
|
||||
),
|
||||
"mechanical_power": MetricsTermCfg(
|
||||
func=mdp.mechanical_power,
|
||||
params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))},
|
||||
),
|
||||
}
|
||||
|
||||
@@ -298,6 +317,11 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg:
|
||||
params={"sensor_name": "robot/root_angmom"},
|
||||
),
|
||||
"is_terminated": RewardTermCfg(func=mdp.is_terminated, weight=-200.0),
|
||||
"electrical_power": RewardTermCfg(
|
||||
func=mdp.electrical_power_cost,
|
||||
weight=0.0,
|
||||
params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))},
|
||||
),
|
||||
"joint_acc_l2": RewardTermCfg(func=mdp.joint_acc_l2, weight=-2.5e-7),
|
||||
"joint_pos_limits": RewardTermCfg(func=mdp.joint_pos_limits, weight=-10.0),
|
||||
"action_rate_l2": RewardTermCfg(func=mdp.action_rate_l2, weight=-0.05),
|
||||
|
||||
+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()
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
@@ -64,13 +65,7 @@ out.write_bytes(b'onnx')
|
||||
self.assertTrue((trainer_root / "scripts" / "train.py").is_file())
|
||||
self.assertTrue(
|
||||
(
|
||||
trainer_root
|
||||
/ "src"
|
||||
/ "assets"
|
||||
/ "robots"
|
||||
/ "unitree_go2"
|
||||
/ "xmls"
|
||||
/ "go2.xml"
|
||||
trainer_root / "src" / "assets" / "robots" / "unitree_go2" / "xmls" / "go2.xml"
|
||||
).is_file()
|
||||
)
|
||||
|
||||
@@ -92,6 +87,15 @@ out.write_bytes(b'onnx')
|
||||
)
|
||||
self.assertEqual(command[-2:], ["--gpu-ids", "[0,2]"])
|
||||
|
||||
def test_resolves_reward_preset_to_inline_validated_trainer_argument(self):
|
||||
preset_id = "f" * 32
|
||||
reward_config = {"weights": {"pose": 1.2}, "params": {}}
|
||||
self.manager.preset_resolver = lambda value: reward_config if value == preset_id else None
|
||||
config = self.manager.parse_config(self.payload(rewardPresetId=preset_id))
|
||||
command = self.manager.command_for(config)
|
||||
index = command.index("--reward-config-json")
|
||||
self.assertEqual(json.loads(command[index + 1]), reward_config)
|
||||
|
||||
def test_requires_local_host_origin_and_bearer_token(self):
|
||||
handler = object.__new__(TrainingRequestHandler)
|
||||
handler.access_token = "secret-token-1234"
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
import math
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from tuning.advisor import AdvisorConfig, DeepSeekAdvisor # noqa: E402
|
||||
from tuning.schema import ( # noqa: E402
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
RewardConfigError,
|
||||
merge_proposal,
|
||||
validate_configuration,
|
||||
validate_proposal,
|
||||
)
|
||||
from tuning.scoring import ( # noqa: E402
|
||||
DEFAULT_OBJECTIVE_WEIGHTS,
|
||||
EvaluationError,
|
||||
score_evaluation,
|
||||
)
|
||||
from tuning.storage import TuningStorage # noqa: E402
|
||||
|
||||
|
||||
class RewardSchemaTest(unittest.TestCase):
|
||||
def test_baseline_is_complete_and_energy_is_disabled(self):
|
||||
config = validate_configuration(BASE_REWARD_CONFIGURATION)
|
||||
self.assertEqual(len(config["weights"]), 16)
|
||||
self.assertEqual(config["weights"]["electrical_power"], 0.0)
|
||||
|
||||
def test_sparse_proposal_constraints(self):
|
||||
patch = validate_proposal(
|
||||
{
|
||||
"weights": {"track_linear_velocity": 1.2, "foot_slip": -0.3},
|
||||
"params": {"foot_gait.period": 0.65},
|
||||
},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
)
|
||||
merged = merge_proposal(BASE_REWARD_CONFIGURATION, patch)
|
||||
self.assertEqual(merged["weights"]["track_linear_velocity"], 1.2)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal({"weights": {"track_linear_velocity": 0}}, BASE_REWARD_CONFIGURATION)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal({"weights": {"foot_slip": 0.2}}, BASE_REWARD_CONFIGURATION)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal(
|
||||
{"weights": {"track_linear_velocity": 3.0}}, BASE_REWARD_CONFIGURATION
|
||||
)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal({"params": {"foot_gait.period": math.nan}}, BASE_REWARD_CONFIGURATION)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal(
|
||||
{
|
||||
"weights": {
|
||||
"pose": 1.1,
|
||||
"foot_gait": 0.6,
|
||||
"foot_slip": -0.3,
|
||||
"soft_landing": -0.002,
|
||||
},
|
||||
"params": {"foot_gait.period": 0.65},
|
||||
},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
)
|
||||
|
||||
def test_cross_parameter_order(self):
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal(
|
||||
{"params": {"pose.walking_threshold": 0.5, "pose.running_threshold": 0.4}},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
)
|
||||
|
||||
|
||||
class ScoringTest(unittest.TestCase):
|
||||
baseline = {
|
||||
"linear_velocity_rmse": 0.3,
|
||||
"angular_velocity_rmse": 0.2,
|
||||
"mean_action_acc": 0.1,
|
||||
"orientation_error": 0.2,
|
||||
"fall_rate": 0.1,
|
||||
"slip_velocity": 0.2,
|
||||
"mechanical_power": 100.0,
|
||||
}
|
||||
|
||||
def test_improvement_and_safety_gate(self):
|
||||
better = {key: value * 0.8 for key, value in self.baseline.items()}
|
||||
scored = score_evaluation(self.baseline, better, DEFAULT_OBJECTIVE_WEIGHTS)
|
||||
self.assertTrue(scored["eligible"])
|
||||
self.assertGreater(scored["score"], 0)
|
||||
unsafe = dict(better, fall_rate=0.2)
|
||||
scored = score_evaluation(self.baseline, unsafe)
|
||||
self.assertFalse(scored["eligible"])
|
||||
self.assertEqual(scored["score"], -1.0)
|
||||
|
||||
def test_rejects_missing_and_nonfinite_metrics(self):
|
||||
with self.assertRaises(EvaluationError):
|
||||
score_evaluation(self.baseline, {"fall_rate": 0.1})
|
||||
bad = dict(self.baseline, mechanical_power=math.inf)
|
||||
with self.assertRaises(EvaluationError):
|
||||
score_evaluation(self.baseline, bad)
|
||||
|
||||
|
||||
class AdvisorTest(unittest.TestCase):
|
||||
class Output:
|
||||
weights = {"pose": 1.1}
|
||||
params = {}
|
||||
rationale = "improve posture"
|
||||
expected_impact = {"posture": "better"}
|
||||
confidence = 0.7
|
||||
|
||||
class Result:
|
||||
output = None
|
||||
|
||||
@staticmethod
|
||||
def usage():
|
||||
return type("Usage", (), {"requests": 1, "input_tokens": 10, "output_tokens": 5})()
|
||||
|
||||
class Agent:
|
||||
def __init__(self, output):
|
||||
self.output = output
|
||||
|
||||
def run_sync(self, _prompt):
|
||||
result = AdvisorTest.Result()
|
||||
result.output = self.output
|
||||
return result
|
||||
|
||||
def test_structured_result_is_revalidated_locally(self):
|
||||
advisor = DeepSeekAdvisor(AdvisorConfig("fake"))
|
||||
advisor._cached_agent = self.Agent(self.Output())
|
||||
proposal = advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION)
|
||||
self.assertEqual(proposal["patch"]["weights"]["pose"], 1.1)
|
||||
self.assertEqual(proposal["usage"]["input_tokens"], 10)
|
||||
|
||||
def test_invalid_model_patch_is_rejected(self):
|
||||
output = self.Output()
|
||||
output.weights = {"track_linear_velocity": -1.0}
|
||||
advisor = DeepSeekAdvisor(AdvisorConfig("fake"))
|
||||
advisor._cached_agent = self.Agent(output)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION)
|
||||
|
||||
|
||||
class StorageTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.temporary = tempfile.TemporaryDirectory()
|
||||
self.storage = TuningStorage(Path(self.temporary.name) / "state.sqlite3")
|
||||
|
||||
def tearDown(self):
|
||||
self.temporary.cleanup()
|
||||
|
||||
def test_persists_session_trial_proposal_and_metrics(self):
|
||||
session = self.storage.create_session(
|
||||
"approval", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, False
|
||||
)
|
||||
trial = self.storage.create_trial(
|
||||
session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0"
|
||||
)
|
||||
proposal = self.storage.create_proposal(
|
||||
session["id"],
|
||||
trial["id"],
|
||||
{"weights": {"pose": 1.1}, "params": {}},
|
||||
"test",
|
||||
{},
|
||||
0.8,
|
||||
)
|
||||
self.assertTrue(self.storage.decide_proposal(proposal["id"], "approved", None))
|
||||
self.assertFalse(self.storage.decide_proposal(proposal["id"], "approved", None))
|
||||
points = [
|
||||
("Train/reward", step, float(step), 50.0 if step == 50 else float(step % 7))
|
||||
for step in range(100)
|
||||
]
|
||||
self.storage.insert_metrics(trial["id"], points)
|
||||
sampled = self.storage.metrics(trial["id"], max_points=10)[0]["points"]
|
||||
self.assertEqual(len(sampled), 10)
|
||||
self.assertEqual(sampled[0]["step"], 0)
|
||||
self.assertEqual(sampled[-1]["step"], 99)
|
||||
self.assertIn(50.0, [point["value"] for point in sampled])
|
||||
reopened = TuningStorage(self.storage.path)
|
||||
self.assertEqual(reopened.get_session(session["id"])["mode"], "approval")
|
||||
|
||||
def test_recovery_marks_inflight_records(self):
|
||||
session = self.storage.create_session(
|
||||
"automatic", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, True
|
||||
)
|
||||
trial = self.storage.create_trial(
|
||||
session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0"
|
||||
)
|
||||
self.storage.update_session(session["id"], state="running")
|
||||
self.storage.update_trial(trial["id"], state="training")
|
||||
self.storage.recover_interrupted()
|
||||
self.assertEqual(self.storage.get_session(session["id"])["state"], "interrupted")
|
||||
self.assertEqual(self.storage.get_trial(trial["id"])["state"], "interrupted")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,195 @@
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from tuning.manager import TuningManager # noqa: E402
|
||||
from tuning.process import GpuLease # noqa: E402
|
||||
from tuning.schema import BASE_REWARD_CONFIGURATION # noqa: E402
|
||||
from tuning.scoring import score_evaluation # noqa: E402
|
||||
from tuning.storage import now_iso # noqa: E402
|
||||
|
||||
BASE_METRICS = {
|
||||
"linear_velocity_rmse": 0.3,
|
||||
"angular_velocity_rmse": 0.2,
|
||||
"mean_action_acc": 0.1,
|
||||
"orientation_error": 0.2,
|
||||
"fall_rate": 0.1,
|
||||
"slip_velocity": 0.2,
|
||||
"mechanical_power": 100.0,
|
||||
}
|
||||
|
||||
|
||||
class FakeAdvisor:
|
||||
def capability(self):
|
||||
return {
|
||||
"configured": True,
|
||||
"apiKeyConfigured": True,
|
||||
"frameworkInstalled": True,
|
||||
"model": "fake",
|
||||
"baseUrl": "https://example.invalid",
|
||||
}
|
||||
|
||||
def propose(self, _context, previous):
|
||||
value = min(2.4, previous["weights"]["pose"] * 1.05)
|
||||
return {
|
||||
"patch": {"weights": {"pose": value}, "params": {}},
|
||||
"rationale": "fake",
|
||||
"expectedImpact": {},
|
||||
"confidence": 0.8,
|
||||
"promptHash": "abc",
|
||||
"usage": {},
|
||||
"model": "fake",
|
||||
}
|
||||
|
||||
def test_connection(self):
|
||||
return {"ok": True, "model": "fake", "outputType": "fake"}
|
||||
|
||||
|
||||
class FakeTuningManager(TuningManager):
|
||||
def _execute_trial(self, session, trial, resume_checkpoint=None):
|
||||
del resume_checkpoint
|
||||
factor = max(0.5, 1.0 - 0.03 * trial["number"] - 0.01 * trial["rung"])
|
||||
metrics = {key: value * factor for key, value in BASE_METRICS.items()}
|
||||
trials = self.storage.list_trials(session["id"])
|
||||
baseline = next(
|
||||
(item for item in trials if item["number"] == 0 and item["rung"] == 0), None
|
||||
)
|
||||
if baseline and baseline["evaluation"]:
|
||||
scored = score_evaluation(
|
||||
baseline["evaluation"]["metrics"], metrics, session["objectiveWeights"]
|
||||
)
|
||||
else:
|
||||
scored = {"score": 0.0, "eligible": True, "components": {}}
|
||||
root = self._session_root(session["id"])
|
||||
run = root / trial["runDir"]
|
||||
run.mkdir(parents=True, exist_ok=True)
|
||||
(run / "model_1.pt").write_bytes(b"checkpoint")
|
||||
(run / "policy.onnx").write_bytes(b"onnx")
|
||||
evaluation = {"metrics": metrics, "score": scored}
|
||||
self.storage.update_trial(
|
||||
trial["id"],
|
||||
state="completed",
|
||||
started_at=now_iso(),
|
||||
ended_at=now_iso(),
|
||||
message="fake complete",
|
||||
checkpoint_path=str((run / "model_1.pt").relative_to(root)),
|
||||
policy_path=str((run / "policy.onnx").relative_to(root)),
|
||||
evaluation=evaluation,
|
||||
score=scored["score"],
|
||||
eligible=scored["eligible"],
|
||||
)
|
||||
return self.storage.get_trial(trial["id"])
|
||||
|
||||
|
||||
class TuningManagerTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.temporary = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self.temporary.name)
|
||||
(self.root / "trainer" / "scripts").mkdir(parents=True)
|
||||
(self.root / "trainer" / "scripts" / "evaluate.py").write_text("", encoding="utf-8")
|
||||
self.manager = FakeTuningManager(
|
||||
self.root / "trainer",
|
||||
sys.executable,
|
||||
self.root / "data",
|
||||
GpuLease(),
|
||||
advisor=FakeAdvisor(),
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
self.manager.shutdown()
|
||||
self.temporary.cleanup()
|
||||
|
||||
@staticmethod
|
||||
def payload(mode="automatic"):
|
||||
return {
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"mode": mode,
|
||||
"runName": "test",
|
||||
"numEnvs": 16,
|
||||
"gpuIds": [0],
|
||||
"trialCount": 4,
|
||||
"initialIterations": 1,
|
||||
"middleIterations": 2,
|
||||
"finalIterations": 3,
|
||||
"evalNumEnvs": 8,
|
||||
"evalSteps": 10,
|
||||
}
|
||||
|
||||
def wait_terminal(self, session_id, timeout=5):
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
session = self.manager.detail(session_id)
|
||||
if session["state"] in {"succeeded", "failed", "cancelled"}:
|
||||
return session
|
||||
time.sleep(0.01)
|
||||
self.fail("session did not finish")
|
||||
|
||||
def test_automatic_session_runs_rungs_and_persists_best_artifact(self):
|
||||
session = self.manager.create(self.payload())
|
||||
completed = self.wait_terminal(session["id"])
|
||||
self.assertEqual(completed["state"], "succeeded", completed["message"])
|
||||
self.assertGreaterEqual(len(completed["trials"]), 7)
|
||||
self.assertTrue(self.manager.best_artifact(session["id"]).is_file())
|
||||
self.assertEqual(len(self.manager.storage.list_presets()), 1)
|
||||
|
||||
def test_approval_session_waits_and_accepts_modified_patch(self):
|
||||
session = self.manager.create(self.payload("approval"))
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if detail["state"] == "awaiting_approval":
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("session did not wait for approval")
|
||||
proposal = detail["proposals"][-1]
|
||||
patch = {"weights": {"pose": 1.1}, "params": {}}
|
||||
approved = self.manager.approve(
|
||||
session["id"], proposal["id"], {"feedback": "ok", "patch": patch}
|
||||
)
|
||||
self.assertEqual(approved["proposals"][-1]["state"], "approved")
|
||||
self.manager.cancel(session["id"])
|
||||
self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled")
|
||||
|
||||
def test_resume_discards_only_interrupted_trial_and_continues(self):
|
||||
mode, config, objective, fallback = self.manager.parse_create(self.payload())
|
||||
session = self.manager.storage.create_session(mode, config, objective, fallback)
|
||||
baseline = self.manager.storage.create_trial(
|
||||
session["id"],
|
||||
0,
|
||||
0,
|
||||
1,
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
None,
|
||||
"trial-000-rung-0",
|
||||
)
|
||||
self.manager._execute_trial(self.manager.storage.get_session(session["id"]), baseline)
|
||||
interrupted = self.manager.storage.create_trial(
|
||||
session["id"],
|
||||
1,
|
||||
0,
|
||||
1,
|
||||
baseline["rewardConfig"],
|
||||
None,
|
||||
"trial-001-rung-0",
|
||||
)
|
||||
self.manager.storage.update_trial(interrupted["id"], state="interrupted")
|
||||
self.manager.storage.update_session(session["id"], state="interrupted")
|
||||
self.manager.resume(session["id"])
|
||||
completed = self.wait_terminal(session["id"])
|
||||
self.assertEqual(completed["state"], "succeeded", completed["message"])
|
||||
self.assertNotIn(interrupted["id"], [trial["id"] for trial in completed["trials"]])
|
||||
|
||||
def test_create_validation_and_agent_capability(self):
|
||||
self.assertTrue(self.manager.capability()["configured"])
|
||||
with self.assertRaisesRegex(Exception, "只支持"):
|
||||
self.manager.parse_create({"taskId": "Other"})
|
||||
self.assertEqual(self.manager.test_agent()["model"], "fake")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,26 @@
|
||||
"""Reward auto-tuning support for the local training service."""
|
||||
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
PARAMETER_SPECS,
|
||||
WEIGHT_SPECS,
|
||||
RewardConfigError,
|
||||
apply_reward_configuration,
|
||||
merge_proposal,
|
||||
validate_configuration,
|
||||
validate_proposal,
|
||||
)
|
||||
from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation
|
||||
|
||||
__all__ = [
|
||||
"BASE_REWARD_CONFIGURATION",
|
||||
"DEFAULT_OBJECTIVE_WEIGHTS",
|
||||
"PARAMETER_SPECS",
|
||||
"WEIGHT_SPECS",
|
||||
"RewardConfigError",
|
||||
"apply_reward_configuration",
|
||||
"merge_proposal",
|
||||
"score_evaluation",
|
||||
"validate_configuration",
|
||||
"validate_proposal",
|
||||
]
|
||||
@@ -0,0 +1,134 @@
|
||||
"""PydanticAI adapter for the DeepSeek reward-tuning advisor."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from .schema import validate_proposal
|
||||
|
||||
SYSTEM_PROMPT = """你是 Unitree Go2 强化学习奖励调参专家。
|
||||
只根据提供的数值配置、训练曲线摘要和固定评估结果提出下一轮稀疏修改。
|
||||
必须优先保持速度跟踪与跌倒安全门槛;每轮最多修改四个白名单标量,不得改变符号、函数、传感器或结构。
|
||||
不要建议 Python 代码、命令、文件路径或白名单外参数。输出必须符合 RewardProposal schema。
|
||||
"""
|
||||
|
||||
|
||||
class AdvisorUnavailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AdvisorConfig:
|
||||
api_key: str | None
|
||||
base_url: str = "https://api.deepseek.com"
|
||||
model: str = "deepseek-v4-flash"
|
||||
|
||||
@classmethod
|
||||
def from_environment(cls) -> AdvisorConfig:
|
||||
return cls(
|
||||
api_key=os.environ.get("DEEPSEEK_API_KEY"),
|
||||
base_url=os.environ.get("MUJOCO_TUNING_AGENT_BASE_URL", "https://api.deepseek.com"),
|
||||
model=os.environ.get("MUJOCO_TUNING_AGENT_MODEL", "deepseek-v4-flash"),
|
||||
)
|
||||
|
||||
|
||||
class DeepSeekAdvisor:
|
||||
def __init__(self, config: AdvisorConfig | None = None):
|
||||
self.config = config or AdvisorConfig.from_environment()
|
||||
self._cached_agent = None
|
||||
|
||||
def capability(self) -> dict[str, Any]:
|
||||
try:
|
||||
import pydantic_ai # noqa: F401
|
||||
except ImportError:
|
||||
installed = False
|
||||
else:
|
||||
installed = True
|
||||
return {
|
||||
"configured": bool(self.config.api_key) and installed,
|
||||
"apiKeyConfigured": bool(self.config.api_key),
|
||||
"frameworkInstalled": installed,
|
||||
"model": self.config.model,
|
||||
"baseUrl": self.config.base_url,
|
||||
}
|
||||
|
||||
def _agent(self):
|
||||
if self._cached_agent is not None:
|
||||
return self._cached_agent
|
||||
if not self.config.api_key:
|
||||
raise AdvisorUnavailable("未配置 DEEPSEEK_API_KEY")
|
||||
try:
|
||||
import httpx2
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic_ai import Agent, PromptedOutput
|
||||
from pydantic_ai.models.openai import OpenAIChatModel
|
||||
from pydantic_ai.providers.openai import OpenAIProvider
|
||||
except ImportError as error:
|
||||
raise AdvisorUnavailable(
|
||||
"缺少 PydanticAI,请安装 training_server/requirements.txt"
|
||||
) from error
|
||||
|
||||
class RewardProposalOutput(BaseModel):
|
||||
weights: dict[str, float] = Field(default_factory=dict)
|
||||
params: dict[str, float] = Field(default_factory=dict)
|
||||
rationale: str = Field(min_length=1, max_length=2000)
|
||||
expected_impact: dict[str, str] = Field(default_factory=dict)
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
|
||||
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("ALL_PROXY")
|
||||
if proxy and proxy.startswith("socks://"):
|
||||
proxy = "socks5://" + proxy.removeprefix("socks://")
|
||||
http_client = httpx2.AsyncClient(proxy=proxy, trust_env=False, timeout=60.0)
|
||||
provider = OpenAIProvider(
|
||||
base_url=self.config.base_url, api_key=self.config.api_key, http_client=http_client
|
||||
)
|
||||
model = OpenAIChatModel(self.config.model, provider=provider) # type: ignore[arg-type]
|
||||
self._cached_agent = Agent(
|
||||
model,
|
||||
output_type=PromptedOutput(RewardProposalOutput),
|
||||
system_prompt=SYSTEM_PROMPT,
|
||||
retries=2,
|
||||
model_settings={"temperature": 0.2},
|
||||
)
|
||||
return self._cached_agent
|
||||
|
||||
def propose(self, context: dict[str, Any], previous: dict) -> dict[str, Any]:
|
||||
prompt = json.dumps(context, ensure_ascii=False, separators=(",", ":"), allow_nan=False)
|
||||
result = self._agent().run_sync(prompt)
|
||||
output = result.output
|
||||
patch = validate_proposal(
|
||||
{"weights": dict(output.weights), "params": dict(output.params)}, previous
|
||||
)
|
||||
try:
|
||||
usage = result.usage()
|
||||
usage_value = {
|
||||
key: getattr(usage, key)
|
||||
for key in ("requests", "input_tokens", "output_tokens", "total_tokens")
|
||||
if getattr(usage, key, None) is not None
|
||||
}
|
||||
except (AttributeError, TypeError):
|
||||
usage_value = {}
|
||||
return {
|
||||
"patch": patch,
|
||||
"rationale": output.rationale,
|
||||
"expectedImpact": dict(output.expected_impact),
|
||||
"confidence": float(output.confidence),
|
||||
"promptHash": hashlib.sha256(prompt.encode()).hexdigest(),
|
||||
"usage": usage_value,
|
||||
"model": self.config.model,
|
||||
}
|
||||
|
||||
def test_connection(self) -> dict[str, Any]:
|
||||
base = {
|
||||
"weights": {"track_linear_velocity": 1.0},
|
||||
"params": {},
|
||||
"instruction": "仅返回一个合法示例:把 track_linear_velocity 改为 1.1。",
|
||||
}
|
||||
# A minimal full previous config is supplied by callers for actual proposals;
|
||||
# connectivity probing only verifies the provider and structured response path.
|
||||
result = self._agent().run_sync(json.dumps(base, ensure_ascii=False))
|
||||
return {"ok": True, "model": self.config.model, "outputType": type(result.output).__name__}
|
||||
@@ -0,0 +1,715 @@
|
||||
"""Persistent reward tuning session orchestrator."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .advisor import AdvisorUnavailable, DeepSeekAdvisor
|
||||
from .process import GpuLease, ResourceBusyError, terminate_process
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
merge_proposal,
|
||||
validate_proposal,
|
||||
)
|
||||
from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation, validate_objective_weights
|
||||
from .storage import TuningStorage, now_iso
|
||||
from .study import OptunaStudies
|
||||
from .tensorboard import ingest_scalars
|
||||
|
||||
ACTIVE_SESSION_STATES = {
|
||||
"queued",
|
||||
"running",
|
||||
"evaluating",
|
||||
"awaiting_approval",
|
||||
"paused",
|
||||
"interrupted",
|
||||
}
|
||||
RUNNING_STATES = {"queued", "running", "evaluating"}
|
||||
RUN_NAME = re.compile(r"^[A-Za-z0-9_.-]{1,64}$")
|
||||
|
||||
|
||||
class TuningError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class TuningManager:
|
||||
def __init__(
|
||||
self,
|
||||
trainer_root: Path,
|
||||
python: str,
|
||||
data_root: Path,
|
||||
lease: GpuLease,
|
||||
advisor: Any | None = None,
|
||||
):
|
||||
self.trainer_root = trainer_root.expanduser().resolve()
|
||||
self.python = python
|
||||
self.data_root = data_root.expanduser().resolve()
|
||||
self.data_root.mkdir(parents=True, exist_ok=True)
|
||||
self.storage = TuningStorage(self.data_root / "tuning.sqlite3")
|
||||
self.storage.recover_interrupted()
|
||||
self.studies = OptunaStudies(self.data_root)
|
||||
self.lease = lease
|
||||
self.advisor = advisor or DeepSeekAdvisor()
|
||||
self.lock = threading.RLock()
|
||||
self.condition = threading.Condition(self.lock)
|
||||
self.workers: dict[str, threading.Thread] = {}
|
||||
self.processes: dict[str, subprocess.Popen[str]] = {}
|
||||
self.cancel_events: dict[str, threading.Event] = {}
|
||||
|
||||
def capability(self) -> dict[str, Any]:
|
||||
capability = self.advisor.capability()
|
||||
capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()})
|
||||
return capability
|
||||
|
||||
@staticmethod
|
||||
def _integer(payload: dict, name: str, default: int, minimum: int, maximum: int) -> int:
|
||||
value = payload.get(name, default)
|
||||
if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum:
|
||||
raise TuningError(f"{name} 必须在 {minimum}–{maximum} 之间")
|
||||
return value
|
||||
|
||||
def parse_create(self, payload: Any) -> tuple[str, dict, dict, bool]:
|
||||
if not isinstance(payload, dict):
|
||||
raise TuningError("请求体必须是 JSON 对象")
|
||||
mode = payload.get("mode", "automatic")
|
||||
if mode not in ("automatic", "approval"):
|
||||
raise TuningError("mode 必须是 automatic 或 approval")
|
||||
if payload.get("taskId", "Unitree-Go2-Flat") != "Unitree-Go2-Flat":
|
||||
raise TuningError("第一版只支持 Unitree-Go2-Flat")
|
||||
run_name = payload.get("runName", "auto-tune")
|
||||
if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name):
|
||||
raise TuningError("runName 格式无效")
|
||||
gpu_ids = payload.get("gpuIds", [0])
|
||||
if (
|
||||
not isinstance(gpu_ids, list)
|
||||
or not gpu_ids
|
||||
or any(
|
||||
isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 255
|
||||
for value in gpu_ids
|
||||
)
|
||||
):
|
||||
raise TuningError("gpuIds 必须是非空非负整数数组")
|
||||
trial_count = self._integer(payload, "trialCount", 12, 4, 20)
|
||||
rung0 = self._integer(payload, "initialIterations", 300, 1, 1000000)
|
||||
rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000)
|
||||
rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000)
|
||||
config = {
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384),
|
||||
"seed": self._integer(payload, "seed", 42, 0, 2147483647),
|
||||
"runName": run_name,
|
||||
"gpuIds": gpu_ids,
|
||||
"trialCount": trial_count,
|
||||
"rungs": [rung0, rung1, rung2],
|
||||
"promote": [trial_count, min(4, trial_count), min(2, trial_count)],
|
||||
"evalNumEnvs": self._integer(payload, "evalNumEnvs", 256, 1, 4096),
|
||||
"evalSteps": self._integer(payload, "evalSteps", 1000, 10, 100000),
|
||||
"earlyStopPatience": self._integer(payload, "earlyStopPatience", 4, 1, 20),
|
||||
}
|
||||
objective = validate_objective_weights(
|
||||
payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS)
|
||||
)
|
||||
fallback = payload.get("fallbackEnabled", False)
|
||||
if not isinstance(fallback, bool):
|
||||
raise TuningError("fallbackEnabled 必须是布尔值")
|
||||
return mode, config, objective, fallback
|
||||
|
||||
def create(self, payload: Any) -> dict:
|
||||
mode, config, objective, fallback = self.parse_create(payload)
|
||||
capability = self.capability()
|
||||
if not capability["ready"]:
|
||||
raise TuningError("评估入口未就绪")
|
||||
if not capability["configured"] and not fallback:
|
||||
raise TuningError("DeepSeek Agent 未配置;设置 DEEPSEEK_API_KEY 或显式启用 fallback")
|
||||
with self.lock:
|
||||
active = [
|
||||
session
|
||||
for session in self.storage.list_sessions()
|
||||
if session["state"] in ACTIVE_SESSION_STATES
|
||||
]
|
||||
if active:
|
||||
raise ResourceBusyError("已有调参 session 未结束")
|
||||
session = self.storage.create_session(mode, config, objective, fallback)
|
||||
self._start_worker(session["id"], resume=False)
|
||||
return self.detail(session["id"])
|
||||
|
||||
def _start_worker(self, session_id: str, resume: bool) -> None:
|
||||
cancel = threading.Event()
|
||||
self.cancel_events[session_id] = cancel
|
||||
worker = threading.Thread(
|
||||
target=self._run_session,
|
||||
args=(session_id, resume, cancel),
|
||||
name=f"tuning-{session_id[:8]}",
|
||||
daemon=True,
|
||||
)
|
||||
self.workers[session_id] = worker
|
||||
worker.start()
|
||||
|
||||
def detail(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
session["trials"] = self.storage.list_trials(session_id)
|
||||
session["proposals"] = self.storage.list_proposals(session_id)
|
||||
session["audit"] = self.storage.audit_events(session_id)
|
||||
return session
|
||||
|
||||
def list(self) -> list[dict]:
|
||||
return self.storage.list_sessions()
|
||||
|
||||
def _session_root(self, session_id: str) -> Path:
|
||||
root = (self.data_root / "sessions" / session_id).resolve()
|
||||
if not root.is_relative_to(self.data_root):
|
||||
raise TuningError("非法 session 路径")
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
|
||||
def _run_command(
|
||||
self,
|
||||
session_id: str,
|
||||
command: list[str],
|
||||
cwd: Path,
|
||||
environment: dict[str, str],
|
||||
log_path: Path,
|
||||
) -> int:
|
||||
owner = f"tuning:{session_id}"
|
||||
self.lease.acquire(owner)
|
||||
try:
|
||||
process = subprocess.Popen(
|
||||
command,
|
||||
cwd=cwd,
|
||||
env=environment,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
bufsize=1,
|
||||
start_new_session=True,
|
||||
)
|
||||
with self.lock:
|
||||
self.processes[session_id] = process
|
||||
assert process.stdout is not None
|
||||
with log_path.open("a", encoding="utf-8") as log:
|
||||
for line in process.stdout:
|
||||
log.write(line)
|
||||
log.flush()
|
||||
if self.cancel_events[session_id].is_set():
|
||||
terminate_process(process)
|
||||
break
|
||||
return process.wait()
|
||||
finally:
|
||||
with self.lock:
|
||||
self.processes.pop(session_id, None)
|
||||
self.lease.release(owner)
|
||||
|
||||
@staticmethod
|
||||
def _latest_checkpoint(run_dir: Path) -> Path | None:
|
||||
values = []
|
||||
for path in run_dir.glob("model_*.pt"):
|
||||
match = re.fullmatch(r"model_(\d+)\.pt", path.name)
|
||||
if match:
|
||||
values.append((int(match.group(1)), path))
|
||||
return max(values, default=(0, None), key=lambda value: value[0])[1]
|
||||
|
||||
def _execute_trial(
|
||||
self, session: dict, trial: dict, resume_checkpoint: Path | None = None
|
||||
) -> dict:
|
||||
session_id, trial_id = session["id"], trial["id"]
|
||||
root = self._session_root(session_id)
|
||||
run_dir = (root / trial["runDir"]).resolve()
|
||||
if not run_dir.is_relative_to(root):
|
||||
raise TuningError("trial 目录越界")
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
reward_path = run_dir / "reward_config.json"
|
||||
reward_path.write_text(
|
||||
json.dumps(trial["rewardConfig"], ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||
)
|
||||
config = session["config"]
|
||||
command = [
|
||||
self.python,
|
||||
"-u",
|
||||
"scripts/train.py",
|
||||
config["taskId"],
|
||||
f"--env.scene.num-envs={config['numEnvs']}",
|
||||
f"--agent.max-iterations={trial['targetIterations']}",
|
||||
f"--agent.seed={config['seed']}",
|
||||
f"--agent.run-name={config['runName']}-t{trial['number']}-r{trial['rung']}",
|
||||
"--agent.logger=tensorboard",
|
||||
"--agent.upload-model=False",
|
||||
"--gpu-ids",
|
||||
json.dumps(config["gpuIds"], separators=(",", ":")),
|
||||
"--output-dir",
|
||||
str(run_dir),
|
||||
"--reward-config",
|
||||
str(reward_path),
|
||||
]
|
||||
if resume_checkpoint is not None:
|
||||
command.extend(("--resume-checkpoint", str(resume_checkpoint)))
|
||||
environment = os.environ.copy()
|
||||
environment["WANDB_MODE"] = "disabled"
|
||||
environment["WANDB_SILENT"] = "true"
|
||||
self.storage.update_trial(
|
||||
trial_id, state="training", started_at=now_iso(), message="正在训练"
|
||||
)
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="running",
|
||||
message=f"正在训练 trial {trial['number']} / rung {trial['rung']}",
|
||||
)
|
||||
return_code = self._run_command(
|
||||
session_id, command, self.trainer_root, environment, run_dir / "train.log"
|
||||
)
|
||||
ingest_scalars(self.storage, trial_id, run_dir)
|
||||
if self.cancel_events[session_id].is_set():
|
||||
raise TuningError("session 已取消")
|
||||
checkpoint = self._latest_checkpoint(run_dir)
|
||||
policy = run_dir / "policy.onnx"
|
||||
if return_code != 0 or checkpoint is None or not policy.is_file():
|
||||
raise TuningError(f"训练失败(返回码 {return_code})或缺少 checkpoint/policy.onnx")
|
||||
self._wait_if_paused(session_id, self.cancel_events[session_id])
|
||||
|
||||
eval_output = run_dir / "evaluation.json"
|
||||
eval_command = [
|
||||
self.python,
|
||||
"-u",
|
||||
"scripts/evaluate.py",
|
||||
config["taskId"],
|
||||
"--checkpoint",
|
||||
str(checkpoint),
|
||||
"--output",
|
||||
str(eval_output),
|
||||
"--reward-config",
|
||||
str(reward_path),
|
||||
f"--num-envs={config['evalNumEnvs']}",
|
||||
f"--steps-per-seed={config['evalSteps']}",
|
||||
"--gpu-ids",
|
||||
json.dumps(config["gpuIds"], separators=(",", ":")),
|
||||
]
|
||||
self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估")
|
||||
self.storage.update_session(
|
||||
session_id, state="evaluating", message=f"正在评估 trial {trial['number']}"
|
||||
)
|
||||
return_code = self._run_command(
|
||||
session_id, eval_command, self.trainer_root, environment, run_dir / "evaluate.log"
|
||||
)
|
||||
ingest_scalars(self.storage, trial_id, run_dir / "evaluation-events")
|
||||
if return_code != 0 or not eval_output.is_file():
|
||||
raise TuningError(f"评估失败(返回码 {return_code})")
|
||||
evaluation = json.loads(eval_output.read_text(encoding="utf-8"))
|
||||
baseline_trial = self.storage.list_trials(session_id)[0]
|
||||
if baseline_trial["evaluation"] is None:
|
||||
scored = {
|
||||
"eligible": True,
|
||||
"score": 0.0,
|
||||
"components": {},
|
||||
"metrics": evaluation["metrics"],
|
||||
}
|
||||
else:
|
||||
scored = score_evaluation(
|
||||
baseline_trial["evaluation"]["metrics"],
|
||||
evaluation["metrics"],
|
||||
session["objectiveWeights"],
|
||||
)
|
||||
evaluation["score"] = scored
|
||||
rel_checkpoint = str(checkpoint.relative_to(root))
|
||||
rel_policy = str(policy.relative_to(root))
|
||||
self.storage.update_trial(
|
||||
trial_id,
|
||||
state="completed",
|
||||
ended_at=now_iso(),
|
||||
message="训练与评估完成",
|
||||
checkpoint_path=rel_checkpoint,
|
||||
policy_path=rel_policy,
|
||||
evaluation=evaluation,
|
||||
score=scored["score"],
|
||||
eligible=scored["eligible"],
|
||||
)
|
||||
try:
|
||||
optuna_number = self.studies.record(
|
||||
session_id,
|
||||
trial["rewardConfig"],
|
||||
scored["score"],
|
||||
scored["eligible"],
|
||||
trial["rung"],
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id, "optuna_trial_recorded", {"trialId": trial_id, "number": optuna_number}
|
||||
)
|
||||
except Exception as error:
|
||||
self.storage.audit(
|
||||
session_id, "optuna_record_failed", {"trialId": trial_id, "error": str(error)[:500]}
|
||||
)
|
||||
return self.storage.get_trial(trial_id)
|
||||
|
||||
def _best(self, session_id: str, rung: int | None = None) -> dict | None:
|
||||
trials = [
|
||||
trial
|
||||
for trial in self.storage.list_trials(session_id)
|
||||
if trial["state"] == "completed" and trial["eligible"]
|
||||
]
|
||||
if rung is not None:
|
||||
trials = [trial for trial in trials if trial["rung"] == rung]
|
||||
return max(
|
||||
trials,
|
||||
key=lambda trial: trial["score"] if trial["score"] is not None else -999,
|
||||
default=None,
|
||||
)
|
||||
|
||||
def _proposal_context(self, session: dict) -> dict:
|
||||
trials = self.storage.list_trials(session["id"])[-12:]
|
||||
rejected_feedback = [
|
||||
proposal["feedback"]
|
||||
for proposal in self.storage.list_proposals(session["id"])
|
||||
if proposal["state"] == "rejected" and proposal["feedback"]
|
||||
][-4:]
|
||||
return {
|
||||
"task": session["config"]["taskId"],
|
||||
"objectiveWeights": session["objectiveWeights"],
|
||||
"allowlist": "服务端将验证固定 schema;最多四项修改",
|
||||
"rejectedFeedback": rejected_feedback,
|
||||
"trials": [
|
||||
{
|
||||
"number": t["number"],
|
||||
"rung": t["rung"],
|
||||
"score": t["score"],
|
||||
"eligible": t["eligible"],
|
||||
"rewardConfig": t["rewardConfig"],
|
||||
"evaluation": t["evaluation"] and t["evaluation"].get("metrics"),
|
||||
}
|
||||
for t in trials
|
||||
],
|
||||
}
|
||||
|
||||
def _fallback_patch(self, previous: dict, index: int) -> dict:
|
||||
names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip")
|
||||
name = names[index % len(names)]
|
||||
old = previous["weights"][name]
|
||||
factor = 1.1 if index % 2 == 0 else 0.9
|
||||
return validate_proposal({"weights": {name: old * factor}}, previous)
|
||||
|
||||
def _request_proposal(
|
||||
self, session: dict, previous: dict, base_trial_id: str, index: int
|
||||
) -> dict:
|
||||
try:
|
||||
result = self.advisor.propose(self._proposal_context(session), previous)
|
||||
source = "agent"
|
||||
except Exception as error:
|
||||
if not session["fallbackEnabled"]:
|
||||
raise AdvisorUnavailable(str(error)) from error
|
||||
result = {
|
||||
"patch": self._fallback_patch(previous, index),
|
||||
"rationale": f"Agent 不可用,显式 fallback:{error}",
|
||||
"expectedImpact": {},
|
||||
"confidence": 0.2,
|
||||
"promptHash": None,
|
||||
"usage": {},
|
||||
"model": "optuna-fallback",
|
||||
}
|
||||
source = "fallback"
|
||||
proposal = self.storage.create_proposal(
|
||||
session["id"],
|
||||
base_trial_id,
|
||||
result["patch"],
|
||||
result["rationale"],
|
||||
result.get("expectedImpact", {}),
|
||||
result["confidence"],
|
||||
source,
|
||||
)
|
||||
self.storage.audit(
|
||||
session["id"],
|
||||
"proposal_created",
|
||||
{
|
||||
"proposalId": proposal["id"],
|
||||
"source": source,
|
||||
"promptHash": result.get("promptHash"),
|
||||
"usage": result.get("usage", {}),
|
||||
"model": result.get("model"),
|
||||
},
|
||||
)
|
||||
return proposal
|
||||
|
||||
def _wait_for_approval(
|
||||
self, session_id: str, proposal_id: str, cancel: threading.Event
|
||||
) -> dict:
|
||||
with self.condition:
|
||||
while not cancel.is_set():
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["state"] != "pending":
|
||||
return proposal
|
||||
self.condition.wait(timeout=1.0)
|
||||
raise TuningError("session 已取消")
|
||||
|
||||
def _run_session(self, session_id: str, resume: bool, cancel: threading.Event) -> None:
|
||||
try:
|
||||
session = self.storage.get_session(session_id)
|
||||
trials = self.storage.list_trials(session_id)
|
||||
if resume:
|
||||
root = self._session_root(session_id)
|
||||
for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]:
|
||||
run_dir = (root / interrupted["runDir"]).resolve()
|
||||
if run_dir.is_relative_to(root):
|
||||
shutil.rmtree(run_dir, ignore_errors=True)
|
||||
self.storage.delete_trial(interrupted["id"])
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"session_resumed",
|
||||
{
|
||||
"discardedInterruptedTrials": [
|
||||
t["id"] for t in trials if t["state"] == "interrupted"
|
||||
]
|
||||
},
|
||||
)
|
||||
trials = self.storage.list_trials(session_id)
|
||||
if not trials:
|
||||
baseline_dir = "trial-000-rung-0"
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
0,
|
||||
0,
|
||||
session["config"]["rungs"][0],
|
||||
deepcopy(BASE_REWARD_CONFIGURATION),
|
||||
None,
|
||||
baseline_dir,
|
||||
)
|
||||
self._execute_trial(session, trial)
|
||||
session = self.storage.get_session(session_id)
|
||||
completed_rung0 = [
|
||||
t
|
||||
for t in self.storage.list_trials(session_id)
|
||||
if t["rung"] == 0 and t["state"] == "completed"
|
||||
]
|
||||
next_number = len({t["number"] for t in completed_rung0})
|
||||
best_score = max((t["score"] or 0.0 for t in completed_rung0), default=0.0)
|
||||
no_improve = 0
|
||||
while (
|
||||
next_number < session["config"]["trialCount"]
|
||||
and no_improve < session["config"]["earlyStopPatience"]
|
||||
):
|
||||
if cancel.is_set():
|
||||
raise TuningError("session 已取消")
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
base = self._best(session_id, rung=0) or completed_rung0[0]
|
||||
proposal = self._request_proposal(
|
||||
session, base["rewardConfig"], base["id"], next_number
|
||||
)
|
||||
if session["mode"] == "approval":
|
||||
self.storage.update_session(
|
||||
session_id, state="awaiting_approval", message="等待批准 Agent 建议"
|
||||
)
|
||||
proposal = self._wait_for_approval(session_id, proposal["id"], cancel)
|
||||
if proposal["state"] == "rejected":
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"proposal_rejected",
|
||||
{"proposalId": proposal["id"], "feedback": proposal["feedback"]},
|
||||
)
|
||||
continue
|
||||
else:
|
||||
self.storage.decide_proposal(proposal["id"], "approved", "自动模式")
|
||||
proposal = self.storage.get_proposal(proposal["id"])
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
reward_config = merge_proposal(base["rewardConfig"], proposal["patch"])
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
next_number,
|
||||
0,
|
||||
session["config"]["rungs"][0],
|
||||
reward_config,
|
||||
proposal["id"],
|
||||
f"trial-{next_number:03d}-rung-0",
|
||||
)
|
||||
result = self._execute_trial(session, trial)
|
||||
if result["eligible"] and (result["score"] or -999) > best_score + 0.01:
|
||||
best_score = result["score"]
|
||||
no_improve = 0
|
||||
else:
|
||||
no_improve += 1
|
||||
self.storage.update_session(session_id, consecutive_no_improve=no_improve)
|
||||
next_number += 1
|
||||
|
||||
# Promote top configurations; each new rung resumes its own previous checkpoint.
|
||||
for rung in (1, 2):
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
previous = [
|
||||
t
|
||||
for t in self.storage.list_trials(session_id)
|
||||
if t["rung"] == rung - 1 and t["state"] == "completed" and t["eligible"]
|
||||
]
|
||||
previous.sort(key=lambda t: t["score"] or -999, reverse=True)
|
||||
promoted_numbers = {
|
||||
trial["number"]
|
||||
for trial in self.storage.list_trials(session_id)
|
||||
if trial["rung"] == rung and trial["state"] == "completed"
|
||||
}
|
||||
for parent in previous[: session["config"]["promote"][rung]]:
|
||||
if parent["number"] in promoted_numbers:
|
||||
continue
|
||||
if cancel.is_set():
|
||||
raise TuningError("session 已取消")
|
||||
root = self._session_root(session_id)
|
||||
checkpoint = root / parent["checkpointPath"]
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
parent["number"],
|
||||
rung,
|
||||
session["config"]["rungs"][rung],
|
||||
parent["rewardConfig"],
|
||||
parent["proposalId"],
|
||||
f"trial-{parent['number']:03d}-rung-{rung}",
|
||||
)
|
||||
self._execute_trial(session, trial, checkpoint)
|
||||
|
||||
best = (
|
||||
self._best(session_id, rung=2)
|
||||
or self._best(session_id, rung=1)
|
||||
or self._best(session_id, rung=0)
|
||||
)
|
||||
if best is None:
|
||||
raise TuningError("没有通过安全门槛的 trial")
|
||||
preset_name = f"{session['config']['runName']}-{session_id[:8]}"
|
||||
self.storage.save_preset(preset_name, session_id, best["id"], best["rewardConfig"])
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="succeeded",
|
||||
best_trial_id=best["id"],
|
||||
current_trial_id=None,
|
||||
message="调参完成",
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id, "session_completed", {"bestTrialId": best["id"], "preset": preset_name}
|
||||
)
|
||||
except Exception as error:
|
||||
state = self.storage.get_session(session_id)["state"]
|
||||
if cancel.is_set() or state == "cancelled":
|
||||
self.storage.update_session(
|
||||
session_id, state="cancelled", message="调参已取消", current_trial_id=None
|
||||
)
|
||||
else:
|
||||
self.storage.update_session(
|
||||
session_id, state="failed", message=str(error), current_trial_id=None
|
||||
)
|
||||
self.storage.audit(session_id, "session_failed", {"error": str(error)[:1000]})
|
||||
finally:
|
||||
with self.lock:
|
||||
self.workers.pop(session_id, None)
|
||||
self.processes.pop(session_id, None)
|
||||
|
||||
def _wait_if_paused(self, session_id: str, cancel: threading.Event) -> None:
|
||||
with self.condition:
|
||||
while self.storage.get_session(session_id)["state"] == "paused" and not cancel.is_set():
|
||||
self.condition.wait(timeout=1.0)
|
||||
|
||||
def approve(self, session_id: str, proposal_id: str, payload: Any) -> dict:
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["sessionId"] != session_id:
|
||||
raise TuningError("proposal 不属于该 session")
|
||||
patch = proposal["patch"]
|
||||
feedback = None
|
||||
if isinstance(payload, dict):
|
||||
feedback = payload.get("feedback")
|
||||
if "patch" in payload:
|
||||
base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
patch = validate_proposal(payload["patch"], base["rewardConfig"])
|
||||
if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch):
|
||||
raise TuningError("proposal 已处理")
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"proposal_approved",
|
||||
{"proposalId": proposal_id, "modified": patch != proposal["patch"]},
|
||||
)
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def reject(self, session_id: str, proposal_id: str, payload: Any) -> dict:
|
||||
feedback = payload.get("feedback", "") if isinstance(payload, dict) else ""
|
||||
if not isinstance(feedback, str) or len(feedback) > 2000:
|
||||
raise TuningError("feedback 无效")
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["sessionId"] != session_id:
|
||||
raise TuningError("proposal 不属于该 session")
|
||||
if not self.storage.decide_proposal(proposal_id, "rejected", feedback):
|
||||
raise TuningError("proposal 已处理")
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def pause(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in RUNNING_STATES | {"awaiting_approval"}:
|
||||
raise TuningError("当前状态不能暂停")
|
||||
self.storage.update_session(
|
||||
session_id, state="paused", message="已暂停后续调度;当前子进程将完成"
|
||||
)
|
||||
return self.detail(session_id)
|
||||
|
||||
def resume(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] == "paused":
|
||||
self.storage.update_session(session_id, state="running", message="继续调参")
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
elif session["state"] == "interrupted":
|
||||
self.storage.update_session(session_id, state="queued", message="从最近完整结果恢复")
|
||||
self._start_worker(session_id, resume=True)
|
||||
else:
|
||||
raise TuningError("当前状态不能恢复")
|
||||
return self.detail(session_id)
|
||||
|
||||
def cancel(self, session_id: str) -> dict:
|
||||
self.storage.get_session(session_id)
|
||||
self.storage.update_session(session_id, state="cancelled", message="正在取消")
|
||||
event = self.cancel_events.get(session_id)
|
||||
if event:
|
||||
event.set()
|
||||
process = self.processes.get(session_id)
|
||||
if process:
|
||||
terminate_process(process)
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def metrics(
|
||||
self, session_id: str, trial_id: str, tags: list[str] | None, max_points: int
|
||||
) -> dict:
|
||||
trial = self.storage.get_trial(trial_id)
|
||||
if trial["sessionId"] != session_id:
|
||||
raise TuningError("trial 不属于该 session")
|
||||
return {"trialId": trial_id, "series": self.storage.metrics(trial_id, tags, max_points)}
|
||||
|
||||
def best_artifact(self, session_id: str) -> Path:
|
||||
session = self.storage.get_session(session_id)
|
||||
if not session["bestTrialId"]:
|
||||
raise TuningError("尚无最佳策略")
|
||||
trial = self.storage.get_trial(session["bestTrialId"])
|
||||
if not trial["policyPath"]:
|
||||
raise TuningError("最佳策略文件不存在")
|
||||
root = self._session_root(session_id)
|
||||
path = (root / trial["policyPath"]).resolve()
|
||||
if not path.is_relative_to(root) or not path.is_file():
|
||||
raise TuningError("最佳策略文件不存在")
|
||||
return path
|
||||
|
||||
def preset_config(self, preset_id: str) -> dict:
|
||||
return self.storage.get_preset(preset_id)["rewardConfig"]
|
||||
|
||||
def test_agent(self) -> dict:
|
||||
try:
|
||||
return self.advisor.test_connection()
|
||||
except Exception as error:
|
||||
raise TuningError(f"Agent 连接测试失败:{error}") from error
|
||||
|
||||
def shutdown(self) -> None:
|
||||
for session_id in list(self.workers):
|
||||
with suppress(KeyError, TuningError):
|
||||
self.cancel(session_id)
|
||||
for worker in list(self.workers.values()):
|
||||
worker.join(timeout=7)
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Shared GPU lease and process-group lifecycle helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
|
||||
|
||||
class ResourceBusyError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class GpuLease:
|
||||
def __init__(self):
|
||||
self.lock = threading.RLock()
|
||||
self.owner: str | None = None
|
||||
|
||||
def acquire(self, owner: str) -> None:
|
||||
with self.lock:
|
||||
if self.owner is not None and self.owner != owner:
|
||||
raise ResourceBusyError(f"计算资源正由 {self.owner} 使用")
|
||||
self.owner = owner
|
||||
|
||||
def release(self, owner: str) -> None:
|
||||
with self.lock:
|
||||
if self.owner == owner:
|
||||
self.owner = None
|
||||
|
||||
def public(self) -> str | None:
|
||||
with self.lock:
|
||||
return self.owner
|
||||
|
||||
|
||||
def terminate_process(process: subprocess.Popen[str], grace_seconds: float = 5.0) -> None:
|
||||
if process.poll() is not None:
|
||||
return
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGTERM)
|
||||
try:
|
||||
process.wait(timeout=grace_seconds)
|
||||
except subprocess.TimeoutExpired:
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGKILL)
|
||||
@@ -0,0 +1,213 @@
|
||||
"""Pure-Python reward tuning schema shared by the service and trainer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
MAX_PROPOSAL_CHANGES = 4
|
||||
MIN_CHANGE_RATIO = 0.5
|
||||
MAX_CHANGE_RATIO = 2.0
|
||||
|
||||
|
||||
class RewardConfigError(ValueError):
|
||||
"""A reward configuration or proposal violated the allowlist."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NumericSpec:
|
||||
minimum: float
|
||||
maximum: float
|
||||
default: float
|
||||
allow_zero: bool = True
|
||||
|
||||
|
||||
WEIGHT_SPECS: dict[str, NumericSpec] = {
|
||||
"track_linear_velocity": NumericSpec(0.5, 3.0, 1.0, False),
|
||||
"track_angular_velocity": NumericSpec(0.25, 2.0, 1.0, False),
|
||||
"body_orientation_l2": NumericSpec(-3.0, -0.1, -1.0, False),
|
||||
"pose": NumericSpec(0.0, 2.5, 1.0),
|
||||
"body_ang_vel": NumericSpec(-0.2, 0.0, -0.05),
|
||||
"angular_momentum": NumericSpec(-0.1, 0.0, -0.025),
|
||||
"is_terminated": NumericSpec(-400.0, -50.0, -200.0, False),
|
||||
"joint_acc_l2": NumericSpec(-2.0e-6, 0.0, -2.5e-7),
|
||||
"joint_pos_limits": NumericSpec(-30.0, -2.0, -10.0, False),
|
||||
"action_rate_l2": NumericSpec(-0.2, -0.005, -0.05, False),
|
||||
"foot_gait": NumericSpec(0.0, 1.5, 0.5),
|
||||
"foot_clearance": NumericSpec(-3.0, 0.0, -1.0),
|
||||
"foot_slip": NumericSpec(-1.0, 0.0, -0.25),
|
||||
"soft_landing": NumericSpec(-5.0e-3, 0.0, -1.0e-3),
|
||||
"stand_still": NumericSpec(-3.0, 0.0, -1.0),
|
||||
"electrical_power": NumericSpec(-5.0e-3, 0.0, 0.0),
|
||||
}
|
||||
|
||||
PARAMETER_SPECS: dict[str, NumericSpec] = {
|
||||
"track_linear_velocity.std": NumericSpec(0.25, 1.0, math.sqrt(0.25), False),
|
||||
"track_angular_velocity.std": NumericSpec(0.35, 1.2, math.sqrt(0.5), False),
|
||||
"pose.std_standing_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.std_walking_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.std_running_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.walking_threshold": NumericSpec(0.05, 0.5, 0.1, False),
|
||||
"pose.running_threshold": NumericSpec(1.0, 2.5, 1.5, False),
|
||||
"foot_gait.period": NumericSpec(0.4, 0.8, 0.6, False),
|
||||
"foot_gait.threshold": NumericSpec(0.45, 0.65, 0.56, False),
|
||||
"foot_gait.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"foot_clearance.target_height": NumericSpec(0.06, 0.16, 0.1, False),
|
||||
"foot_clearance.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"foot_slip.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"soft_landing.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"stand_still.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
}
|
||||
|
||||
BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = {
|
||||
"weights": {name: spec.default for name, spec in WEIGHT_SPECS.items()},
|
||||
"params": {name: spec.default for name, spec in PARAMETER_SPECS.items()},
|
||||
}
|
||||
|
||||
|
||||
def _number(name: str, value: Any, spec: NumericSpec) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise RewardConfigError(f"{name} 必须是数值")
|
||||
result = float(value)
|
||||
if not math.isfinite(result):
|
||||
raise RewardConfigError(f"{name} 必须是有限数值")
|
||||
if result == 0.0 and not spec.allow_zero:
|
||||
raise RewardConfigError(f"{name} 不允许关闭")
|
||||
if result < spec.minimum or result > spec.maximum:
|
||||
raise RewardConfigError(f"{name} 必须在 {spec.minimum}–{spec.maximum} 之间")
|
||||
return result
|
||||
|
||||
|
||||
def _mapping(value: Any, name: str) -> Mapping[str, Any]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise RewardConfigError(f"{name} 必须是对象")
|
||||
return value
|
||||
|
||||
|
||||
def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None:
|
||||
params = config["params"]
|
||||
if params["pose.walking_threshold"] >= params["pose.running_threshold"]:
|
||||
raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold")
|
||||
|
||||
|
||||
def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a complete configuration and reject missing/unknown fields."""
|
||||
root = _mapping(value, "rewardConfig")
|
||||
if set(root) != {"weights", "params"}:
|
||||
raise RewardConfigError("rewardConfig 只能包含 weights 和 params")
|
||||
raw_weights = _mapping(root["weights"], "weights")
|
||||
raw_params = _mapping(root["params"], "params")
|
||||
if set(raw_weights) != set(WEIGHT_SPECS):
|
||||
raise RewardConfigError("weights 必须完整且不能包含未知奖励项")
|
||||
if set(raw_params) != set(PARAMETER_SPECS):
|
||||
raise RewardConfigError("params 必须完整且不能包含未知参数")
|
||||
config = {
|
||||
"weights": {
|
||||
name: _number(f"weights.{name}", raw_weights[name], spec)
|
||||
for name, spec in WEIGHT_SPECS.items()
|
||||
},
|
||||
"params": {
|
||||
name: _number(f"params.{name}", raw_params[name], spec)
|
||||
for name, spec in PARAMETER_SPECS.items()
|
||||
},
|
||||
}
|
||||
_cross_validate(config)
|
||||
return config
|
||||
|
||||
|
||||
def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config."""
|
||||
current = validate_configuration(previous)
|
||||
root = _mapping(value, "proposal")
|
||||
if not set(root).issubset({"weights", "params"}):
|
||||
raise RewardConfigError("proposal 只能包含 weights 和 params")
|
||||
raw_weights = _mapping(root.get("weights", {}), "weights")
|
||||
raw_params = _mapping(root.get("params", {}), "params")
|
||||
if len(raw_weights) + len(raw_params) == 0:
|
||||
raise RewardConfigError("proposal 至少需要一项修改")
|
||||
if len(raw_weights) + len(raw_params) > MAX_PROPOSAL_CHANGES:
|
||||
raise RewardConfigError(f"proposal 每轮最多修改 {MAX_PROPOSAL_CHANGES} 项")
|
||||
unknown_weights = set(raw_weights) - set(WEIGHT_SPECS)
|
||||
unknown_params = set(raw_params) - set(PARAMETER_SPECS)
|
||||
if unknown_weights:
|
||||
raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}")
|
||||
if unknown_params:
|
||||
raise RewardConfigError(f"未知奖励参数:{', '.join(sorted(unknown_params))}")
|
||||
|
||||
patch: dict[str, dict[str, float]] = {"weights": {}, "params": {}}
|
||||
for name, raw in raw_weights.items():
|
||||
value_number = _number(f"weights.{name}", raw, WEIGHT_SPECS[name])
|
||||
old = current["weights"][name]
|
||||
if old != 0.0 and value_number != 0.0:
|
||||
ratio = abs(value_number / old)
|
||||
if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO:
|
||||
raise RewardConfigError(
|
||||
f"weights.{name} 单轮变化必须在旧值幅度的 "
|
||||
f"{MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×"
|
||||
)
|
||||
if value_number == old:
|
||||
raise RewardConfigError(f"weights.{name} 没有发生变化")
|
||||
patch["weights"][name] = value_number
|
||||
for name, raw in raw_params.items():
|
||||
value_number = _number(f"params.{name}", raw, PARAMETER_SPECS[name])
|
||||
old = current["params"][name]
|
||||
ratio = abs(value_number / old)
|
||||
if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO:
|
||||
raise RewardConfigError(
|
||||
f"params.{name} 单轮变化必须在旧值的 {MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×"
|
||||
)
|
||||
if value_number == old:
|
||||
raise RewardConfigError(f"params.{name} 没有发生变化")
|
||||
patch["params"][name] = value_number
|
||||
|
||||
candidate = deepcopy(current)
|
||||
candidate["weights"].update(patch["weights"])
|
||||
candidate["params"].update(patch["params"])
|
||||
_cross_validate(candidate)
|
||||
return patch
|
||||
|
||||
|
||||
def merge_proposal(previous: Any, proposal: Any) -> dict[str, dict[str, float]]:
|
||||
current = validate_configuration(previous)
|
||||
patch = validate_proposal(proposal, current)
|
||||
merged = deepcopy(current)
|
||||
merged["weights"].update(patch["weights"])
|
||||
merged["params"].update(patch["params"])
|
||||
return validate_configuration(merged)
|
||||
|
||||
|
||||
def apply_reward_configuration(env_cfg: Any, value: Any) -> None:
|
||||
"""Apply a validated full config to a fresh mjlab environment config."""
|
||||
config = validate_configuration(value)
|
||||
for name, weight in config["weights"].items():
|
||||
if name not in env_cfg.rewards:
|
||||
raise RewardConfigError(f"环境缺少奖励项:{name}")
|
||||
env_cfg.rewards[name].weight = weight
|
||||
|
||||
params = config["params"]
|
||||
direct = {
|
||||
"track_linear_velocity.std": ("track_linear_velocity", "std"),
|
||||
"track_angular_velocity.std": ("track_angular_velocity", "std"),
|
||||
"pose.walking_threshold": ("pose", "walking_threshold"),
|
||||
"pose.running_threshold": ("pose", "running_threshold"),
|
||||
"foot_gait.period": ("foot_gait", "period"),
|
||||
"foot_gait.threshold": ("foot_gait", "threshold"),
|
||||
"foot_gait.command_threshold": ("foot_gait", "command_threshold"),
|
||||
"foot_clearance.target_height": ("foot_clearance", "target_height"),
|
||||
"foot_clearance.command_threshold": ("foot_clearance", "command_threshold"),
|
||||
"foot_slip.command_threshold": ("foot_slip", "command_threshold"),
|
||||
"soft_landing.command_threshold": ("soft_landing", "command_threshold"),
|
||||
"stand_still.command_threshold": ("stand_still", "command_threshold"),
|
||||
}
|
||||
for path, (term, parameter) in direct.items():
|
||||
env_cfg.rewards[term].params[parameter] = params[path]
|
||||
for regime in ("standing", "walking", "running"):
|
||||
key = f"std_{regime}"
|
||||
scale = params[f"pose.{key}_scale"]
|
||||
baseline = env_cfg.rewards["pose"].params[key]
|
||||
env_cfg.rewards["pose"].params[key] = {
|
||||
pattern: float(std) * scale for pattern, std in baseline.items()
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
"""Stable, reward-weight-independent evaluation scoring."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
DEFAULT_OBJECTIVE_WEIGHTS = {
|
||||
"velocity_tracking": 0.35,
|
||||
"action_smoothness": 0.20,
|
||||
"posture_stability": 0.15,
|
||||
"fall_avoidance": 0.15,
|
||||
"foot_slip": 0.10,
|
||||
"energy": 0.05,
|
||||
}
|
||||
|
||||
REQUIRED_METRICS = {
|
||||
"linear_velocity_rmse",
|
||||
"angular_velocity_rmse",
|
||||
"mean_action_acc",
|
||||
"orientation_error",
|
||||
"fall_rate",
|
||||
"slip_velocity",
|
||||
"mechanical_power",
|
||||
}
|
||||
|
||||
PHYSICAL_FLOORS = {
|
||||
"linear_velocity_rmse": 0.10,
|
||||
"angular_velocity_rmse": 0.10,
|
||||
"mean_action_acc": 0.01,
|
||||
"orientation_error": 0.05,
|
||||
"fall_rate": 0.02,
|
||||
"slip_velocity": 0.05,
|
||||
"mechanical_power": 10.0,
|
||||
}
|
||||
|
||||
|
||||
class EvaluationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def validate_objective_weights(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, Mapping) or set(value) != set(DEFAULT_OBJECTIVE_WEIGHTS):
|
||||
raise EvaluationError("objectiveWeights 必须完整包含六个目标")
|
||||
result: dict[str, float] = {}
|
||||
for key in DEFAULT_OBJECTIVE_WEIGHTS:
|
||||
raw = value[key]
|
||||
if isinstance(raw, bool) or not isinstance(raw, (int, float)):
|
||||
raise EvaluationError(f"objectiveWeights.{key} 必须是数值")
|
||||
number = float(raw)
|
||||
if not math.isfinite(number) or number < 0.0 or number > 1.0:
|
||||
raise EvaluationError(f"objectiveWeights.{key} 必须在 0–1 之间")
|
||||
result[key] = number
|
||||
if not math.isclose(sum(result.values()), 1.0, abs_tol=1.0e-6):
|
||||
raise EvaluationError("objectiveWeights 总和必须为 1")
|
||||
return result
|
||||
|
||||
|
||||
def validate_metrics(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise EvaluationError("metrics 必须是对象")
|
||||
missing = REQUIRED_METRICS - set(value)
|
||||
if missing:
|
||||
raise EvaluationError(f"metrics 缺少:{', '.join(sorted(missing))}")
|
||||
result: dict[str, float] = {}
|
||||
for key in REQUIRED_METRICS:
|
||||
raw = value[key]
|
||||
if isinstance(raw, bool) or not isinstance(raw, (int, float)):
|
||||
raise EvaluationError(f"metrics.{key} 必须是数值")
|
||||
number = float(raw)
|
||||
if not math.isfinite(number) or number < 0.0:
|
||||
raise EvaluationError(f"metrics.{key} 必须是非负有限数值")
|
||||
result[key] = number
|
||||
if result["fall_rate"] > 1.0:
|
||||
raise EvaluationError("metrics.fall_rate 必须在 0–1 之间")
|
||||
return result
|
||||
|
||||
|
||||
def _improvement(baseline: Mapping[str, float], current: Mapping[str, float], key: str) -> float:
|
||||
scale = max(abs(baseline[key]), PHYSICAL_FLOORS[key])
|
||||
return max(-1.0, min(1.0, (baseline[key] - current[key]) / scale))
|
||||
|
||||
|
||||
def score_evaluation(
|
||||
baseline_value: Any,
|
||||
current_value: Any,
|
||||
objective_weights: Any = DEFAULT_OBJECTIVE_WEIGHTS,
|
||||
) -> dict[str, Any]:
|
||||
baseline = validate_metrics(baseline_value)
|
||||
current = validate_metrics(current_value)
|
||||
weights = validate_objective_weights(objective_weights)
|
||||
components = {
|
||||
"velocity_tracking": 0.8 * _improvement(baseline, current, "linear_velocity_rmse")
|
||||
+ 0.2 * _improvement(baseline, current, "angular_velocity_rmse"),
|
||||
"action_smoothness": _improvement(baseline, current, "mean_action_acc"),
|
||||
"posture_stability": _improvement(baseline, current, "orientation_error"),
|
||||
"fall_avoidance": _improvement(baseline, current, "fall_rate"),
|
||||
"foot_slip": _improvement(baseline, current, "slip_velocity"),
|
||||
"energy": _improvement(baseline, current, "mechanical_power"),
|
||||
}
|
||||
tracking_limit = max(
|
||||
baseline["linear_velocity_rmse"] * 1.05, baseline["linear_velocity_rmse"] + 1.0e-6
|
||||
)
|
||||
eligible = (
|
||||
current["fall_rate"] <= baseline["fall_rate"] + 0.02
|
||||
and current["linear_velocity_rmse"] <= tracking_limit
|
||||
)
|
||||
total = sum(weights[key] * components[key] for key in weights)
|
||||
if not eligible:
|
||||
total = min(total, -1.0)
|
||||
return {
|
||||
"eligible": eligible,
|
||||
"score": total,
|
||||
"components": components,
|
||||
"metrics": current,
|
||||
"baselineMetrics": baseline,
|
||||
}
|
||||
@@ -0,0 +1,521 @@
|
||||
"""SQLite persistence for tuning sessions, trials, proposals and scalar data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
|
||||
|
||||
def now_iso() -> str:
|
||||
return datetime.now(UTC).isoformat()
|
||||
|
||||
|
||||
def _json(value: Any) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False)
|
||||
|
||||
|
||||
def _decode(value: str | None) -> Any:
|
||||
return json.loads(value) if value else None
|
||||
|
||||
|
||||
def _lttb(points: list[dict], threshold: int) -> list[dict]:
|
||||
"""Largest-Triangle-Three-Buckets downsampling preserving peaks and endpoints."""
|
||||
if threshold >= len(points) or threshold < 3:
|
||||
return points[:threshold]
|
||||
sampled = [points[0]]
|
||||
bucket_width = (len(points) - 2) / (threshold - 2)
|
||||
anchor_index = 0
|
||||
for bucket in range(threshold - 2):
|
||||
average_start = int((bucket + 1) * bucket_width) + 1
|
||||
average_end = min(int((bucket + 2) * bucket_width) + 1, len(points))
|
||||
average_bucket = points[average_start:average_end] or [points[-1]]
|
||||
average_x = sum(point["step"] for point in average_bucket) / len(average_bucket)
|
||||
average_y = sum(point["value"] for point in average_bucket) / len(average_bucket)
|
||||
range_start = int(bucket * bucket_width) + 1
|
||||
range_end = min(int((bucket + 1) * bucket_width) + 1, len(points) - 1)
|
||||
anchor = points[anchor_index]
|
||||
selected_index = range_start
|
||||
maximum_area = -1.0
|
||||
for index in range(range_start, max(range_start + 1, range_end)):
|
||||
point = points[index]
|
||||
area = abs(
|
||||
(anchor["step"] - average_x) * (point["value"] - anchor["value"])
|
||||
- (anchor["step"] - point["step"]) * (average_y - anchor["value"])
|
||||
)
|
||||
if area > maximum_area:
|
||||
maximum_area = area
|
||||
selected_index = index
|
||||
sampled.append(points[selected_index])
|
||||
anchor_index = selected_index
|
||||
sampled.append(points[-1])
|
||||
return sampled
|
||||
|
||||
|
||||
class TuningStorage:
|
||||
def __init__(self, path: Path):
|
||||
self.path = path.expanduser().resolve()
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self.local = threading.local()
|
||||
self._migrate()
|
||||
|
||||
def connection(self) -> sqlite3.Connection:
|
||||
connection = getattr(self.local, "connection", None)
|
||||
if connection is None:
|
||||
connection = sqlite3.connect(self.path, timeout=10, isolation_level=None)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys=ON")
|
||||
connection.execute("PRAGMA journal_mode=WAL")
|
||||
connection.execute("PRAGMA busy_timeout=10000")
|
||||
self.local.connection = connection
|
||||
return connection
|
||||
|
||||
@contextmanager
|
||||
def transaction(self) -> Iterator[sqlite3.Connection]:
|
||||
connection = self.connection()
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
yield connection
|
||||
connection.execute("COMMIT")
|
||||
except Exception:
|
||||
connection.execute("ROLLBACK")
|
||||
raise
|
||||
|
||||
def _migrate(self) -> None:
|
||||
connection = self.connection()
|
||||
connection.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations(version INTEGER PRIMARY KEY);
|
||||
CREATE TABLE IF NOT EXISTS sessions(
|
||||
id TEXT PRIMARY KEY, state TEXT NOT NULL, mode TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL, updated_at TEXT NOT NULL,
|
||||
config_json TEXT NOT NULL, objective_json TEXT NOT NULL,
|
||||
message TEXT NOT NULL, current_trial_id TEXT, best_trial_id TEXT,
|
||||
consecutive_no_improve INTEGER NOT NULL DEFAULT 0,
|
||||
fallback_enabled INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS trials(
|
||||
id TEXT PRIMARY KEY,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
number INTEGER NOT NULL, state TEXT NOT NULL, rung INTEGER NOT NULL,
|
||||
target_iterations INTEGER NOT NULL, reward_config_json TEXT NOT NULL,
|
||||
proposal_id TEXT, run_dir TEXT NOT NULL, checkpoint_path TEXT,
|
||||
policy_path TEXT, evaluation_json TEXT, score REAL, eligible INTEGER,
|
||||
created_at TEXT NOT NULL, started_at TEXT, ended_at TEXT, message TEXT NOT NULL,
|
||||
UNIQUE(session_id, number, rung)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS proposals(
|
||||
id TEXT PRIMARY KEY,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
base_trial_id TEXT, state TEXT NOT NULL, source TEXT NOT NULL,
|
||||
patch_json TEXT NOT NULL, rationale TEXT NOT NULL,
|
||||
expected_json TEXT, confidence REAL NOT NULL,
|
||||
created_at TEXT NOT NULL, decided_at TEXT, feedback TEXT
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS metric_points(
|
||||
trial_id TEXT NOT NULL REFERENCES trials(id) ON DELETE CASCADE,
|
||||
tag TEXT NOT NULL, step INTEGER NOT NULL,
|
||||
wall_time REAL NOT NULL, value REAL NOT NULL,
|
||||
PRIMARY KEY(trial_id, tag, step)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS audit_events(
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
event_type TEXT NOT NULL, payload_json TEXT NOT NULL, created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS presets(
|
||||
id TEXT PRIMARY KEY, name TEXT NOT NULL UNIQUE, session_id TEXT NOT NULL,
|
||||
trial_id TEXT NOT NULL, reward_config_json TEXT NOT NULL, created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_trials_session ON trials(session_id, number, rung);
|
||||
CREATE INDEX IF NOT EXISTS idx_proposals_session ON proposals(session_id, created_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_metrics_trial_tag ON metric_points(trial_id, tag, step);
|
||||
"""
|
||||
)
|
||||
connection.execute(
|
||||
"INSERT OR IGNORE INTO schema_migrations(version) VALUES (?)", (SCHEMA_VERSION,)
|
||||
)
|
||||
|
||||
def recover_interrupted(self) -> None:
|
||||
at = now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"UPDATE trials SET state='interrupted', ended_at=?, "
|
||||
"message='服务重启中断,等待显式恢复' "
|
||||
"WHERE state IN ('training','evaluating')",
|
||||
(at,),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET state='interrupted', updated_at=?, "
|
||||
"message='服务重启中断,可从完整 checkpoint 恢复' "
|
||||
"WHERE state IN ('running','evaluating')",
|
||||
(at,),
|
||||
)
|
||||
|
||||
def create_session(self, mode: str, config: dict, objective: dict, fallback: bool) -> dict:
|
||||
session_id, at = uuid.uuid4().hex, now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"INSERT INTO sessions("
|
||||
"id,state,mode,created_at,updated_at,config_json,objective_json,"
|
||||
"message,fallback_enabled) "
|
||||
"VALUES (?, 'queued', ?, ?, ?, ?, ?, '等待基线训练', ?)",
|
||||
(session_id, mode, at, at, _json(config), _json(objective), int(fallback)),
|
||||
)
|
||||
connection.execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
(session_id, "session_created", _json({"mode": mode}), at),
|
||||
)
|
||||
return self.get_session(session_id)
|
||||
|
||||
def _session(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"state": row["state"],
|
||||
"mode": row["mode"],
|
||||
"createdAt": row["created_at"],
|
||||
"updatedAt": row["updated_at"],
|
||||
"config": _decode(row["config_json"]),
|
||||
"objectiveWeights": _decode(row["objective_json"]),
|
||||
"message": row["message"],
|
||||
"currentTrialId": row["current_trial_id"],
|
||||
"bestTrialId": row["best_trial_id"],
|
||||
"consecutiveNoImprove": row["consecutive_no_improve"],
|
||||
"fallbackEnabled": bool(row["fallback_enabled"]),
|
||||
}
|
||||
|
||||
def get_session(self, session_id: str) -> dict:
|
||||
row = (
|
||||
self.connection().execute("SELECT * FROM sessions WHERE id=?", (session_id,)).fetchone()
|
||||
)
|
||||
if row is None:
|
||||
raise KeyError(session_id)
|
||||
return self._session(row)
|
||||
|
||||
def list_sessions(self, limit: int = 50) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM sessions ORDER BY created_at DESC LIMIT ?", (limit,))
|
||||
.fetchall()
|
||||
)
|
||||
return [self._session(row) for row in rows]
|
||||
|
||||
def update_session(self, session_id: str, **changes: Any) -> bool:
|
||||
columns = {
|
||||
"state": "state",
|
||||
"message": "message",
|
||||
"current_trial_id": "current_trial_id",
|
||||
"best_trial_id": "best_trial_id",
|
||||
"consecutive_no_improve": "consecutive_no_improve",
|
||||
}
|
||||
values, assignments = [], []
|
||||
for key, value in changes.items():
|
||||
if key not in columns:
|
||||
raise ValueError(key)
|
||||
assignments.append(f"{columns[key]}=?")
|
||||
values.append(value)
|
||||
assignments.append("updated_at=?")
|
||||
values.extend((now_iso(), session_id))
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE sessions SET {', '.join(assignments)} WHERE id=?", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def create_trial(
|
||||
self,
|
||||
session_id: str,
|
||||
number: int,
|
||||
rung: int,
|
||||
target: int,
|
||||
reward_config: dict,
|
||||
proposal_id: str | None,
|
||||
run_dir: str,
|
||||
) -> dict:
|
||||
trial_id, at = uuid.uuid4().hex, now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"INSERT INTO trials("
|
||||
"id,session_id,number,state,rung,target_iterations,reward_config_json,"
|
||||
"proposal_id,run_dir,created_at,message) "
|
||||
"VALUES (?,?,?,'queued',?,?,?,?,?,?,'等待训练')",
|
||||
(
|
||||
trial_id,
|
||||
session_id,
|
||||
number,
|
||||
rung,
|
||||
target,
|
||||
_json(reward_config),
|
||||
proposal_id,
|
||||
run_dir,
|
||||
at,
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET current_trial_id=?,updated_at=? WHERE id=?",
|
||||
(trial_id, at, session_id),
|
||||
)
|
||||
return self.get_trial(trial_id)
|
||||
|
||||
def _trial(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"sessionId": row["session_id"],
|
||||
"number": row["number"],
|
||||
"state": row["state"],
|
||||
"rung": row["rung"],
|
||||
"targetIterations": row["target_iterations"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"proposalId": row["proposal_id"],
|
||||
"runDir": row["run_dir"],
|
||||
"checkpointPath": row["checkpoint_path"],
|
||||
"policyPath": row["policy_path"],
|
||||
"evaluation": _decode(row["evaluation_json"]),
|
||||
"score": row["score"],
|
||||
"eligible": None if row["eligible"] is None else bool(row["eligible"]),
|
||||
"createdAt": row["created_at"],
|
||||
"startedAt": row["started_at"],
|
||||
"endedAt": row["ended_at"],
|
||||
"message": row["message"],
|
||||
}
|
||||
|
||||
def get_trial(self, trial_id: str) -> dict:
|
||||
row = self.connection().execute("SELECT * FROM trials WHERE id=?", (trial_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(trial_id)
|
||||
return self._trial(row)
|
||||
|
||||
def list_trials(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM trials WHERE session_id=? ORDER BY number,rung", (session_id,))
|
||||
.fetchall()
|
||||
)
|
||||
return [self._trial(row) for row in rows]
|
||||
|
||||
def delete_trial(self, trial_id: str) -> bool:
|
||||
cursor = self.connection().execute(
|
||||
"DELETE FROM trials WHERE id=? AND state='interrupted'", (trial_id,)
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def update_trial(self, trial_id: str, **changes: Any) -> bool:
|
||||
columns = {
|
||||
"state": "state",
|
||||
"message": "message",
|
||||
"checkpoint_path": "checkpoint_path",
|
||||
"policy_path": "policy_path",
|
||||
"score": "score",
|
||||
"eligible": "eligible",
|
||||
"started_at": "started_at",
|
||||
"ended_at": "ended_at",
|
||||
"evaluation": "evaluation_json",
|
||||
}
|
||||
values, assignments = [], []
|
||||
for key, value in changes.items():
|
||||
if key not in columns:
|
||||
raise ValueError(key)
|
||||
if key == "evaluation":
|
||||
value = _json(value)
|
||||
if key == "eligible":
|
||||
value = int(value)
|
||||
assignments.append(f"{columns[key]}=?")
|
||||
values.append(value)
|
||||
values.append(trial_id)
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE trials SET {', '.join(assignments)} WHERE id=?", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def create_proposal(
|
||||
self,
|
||||
session_id: str,
|
||||
base_trial_id: str | None,
|
||||
patch: dict,
|
||||
rationale: str,
|
||||
expected: Any,
|
||||
confidence: float,
|
||||
source: str = "agent",
|
||||
) -> dict:
|
||||
proposal_id, at = uuid.uuid4().hex, now_iso()
|
||||
self.connection().execute(
|
||||
"INSERT INTO proposals("
|
||||
"id,session_id,base_trial_id,state,source,patch_json,rationale,"
|
||||
"expected_json,confidence,created_at) "
|
||||
"VALUES (?,?,?,'pending',?,?,?,?,?,?)",
|
||||
(
|
||||
proposal_id,
|
||||
session_id,
|
||||
base_trial_id,
|
||||
source,
|
||||
_json(patch),
|
||||
rationale,
|
||||
_json(expected),
|
||||
confidence,
|
||||
at,
|
||||
),
|
||||
)
|
||||
return self.get_proposal(proposal_id)
|
||||
|
||||
def _proposal(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"sessionId": row["session_id"],
|
||||
"baseTrialId": row["base_trial_id"],
|
||||
"state": row["state"],
|
||||
"source": row["source"],
|
||||
"patch": _decode(row["patch_json"]),
|
||||
"rationale": row["rationale"],
|
||||
"expectedImpact": _decode(row["expected_json"]),
|
||||
"confidence": row["confidence"],
|
||||
"createdAt": row["created_at"],
|
||||
"decidedAt": row["decided_at"],
|
||||
"feedback": row["feedback"],
|
||||
}
|
||||
|
||||
def get_proposal(self, proposal_id: str) -> dict:
|
||||
row = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM proposals WHERE id=?", (proposal_id,))
|
||||
.fetchone()
|
||||
)
|
||||
if row is None:
|
||||
raise KeyError(proposal_id)
|
||||
return self._proposal(row)
|
||||
|
||||
def list_proposals(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
"SELECT * FROM proposals WHERE session_id=? ORDER BY created_at", (session_id,)
|
||||
)
|
||||
.fetchall()
|
||||
)
|
||||
return [self._proposal(row) for row in rows]
|
||||
|
||||
def decide_proposal(
|
||||
self, proposal_id: str, state: str, feedback: str | None, patch: dict | None = None
|
||||
) -> bool:
|
||||
at = now_iso()
|
||||
assignments, values = ["state=?", "feedback=?", "decided_at=?"], [state, feedback, at]
|
||||
if patch is not None:
|
||||
assignments.append("patch_json=?")
|
||||
values.append(_json(patch))
|
||||
values.extend((proposal_id,))
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE proposals SET {', '.join(assignments)} WHERE id=? AND state='pending'", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def insert_metrics(self, trial_id: str, points: list[tuple[str, int, float, float]]) -> None:
|
||||
self.connection().executemany(
|
||||
"INSERT INTO metric_points(trial_id,tag,step,wall_time,value) "
|
||||
"VALUES (?,?,?,?,?) ON CONFLICT(trial_id,tag,step) DO UPDATE SET "
|
||||
"wall_time=excluded.wall_time,value=excluded.value",
|
||||
[(trial_id, *point) for point in points],
|
||||
)
|
||||
|
||||
def metrics(
|
||||
self, trial_id: str, tags: list[str] | None = None, max_points: int = 1000
|
||||
) -> list[dict]:
|
||||
parameters: list[Any] = [trial_id]
|
||||
clause = "trial_id=?"
|
||||
if tags:
|
||||
clause += f" AND tag IN ({','.join('?' for _ in tags)})"
|
||||
parameters.extend(tags)
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
f"SELECT tag,step,wall_time,value FROM metric_points "
|
||||
f"WHERE {clause} ORDER BY tag,step",
|
||||
parameters,
|
||||
)
|
||||
.fetchall()
|
||||
)
|
||||
grouped: dict[str, list[dict]] = {}
|
||||
for row in rows:
|
||||
grouped.setdefault(row["tag"], []).append(
|
||||
{"step": row["step"], "wallTime": row["wall_time"], "value": row["value"]}
|
||||
)
|
||||
series = []
|
||||
for tag, values in grouped.items():
|
||||
if len(values) > max_points:
|
||||
values = _lttb(values, max_points)
|
||||
series.append({"tag": tag, "points": values})
|
||||
return series
|
||||
|
||||
def audit(self, session_id: str, event_type: str, payload: Any) -> None:
|
||||
self.connection().execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
(session_id, event_type, _json(payload), now_iso()),
|
||||
)
|
||||
|
||||
def audit_events(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM audit_events WHERE session_id=? ORDER BY id", (session_id,))
|
||||
.fetchall()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": row["id"],
|
||||
"type": row["event_type"],
|
||||
"payload": _decode(row["payload_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def save_preset(self, name: str, session_id: str, trial_id: str, reward_config: dict) -> dict:
|
||||
preset_id, at = uuid.uuid4().hex, now_iso()
|
||||
self.connection().execute(
|
||||
"INSERT INTO presets(id,name,session_id,trial_id,reward_config_json,created_at) "
|
||||
"VALUES (?,?,?,?,?,?)",
|
||||
(preset_id, name, session_id, trial_id, _json(reward_config), at),
|
||||
)
|
||||
return {
|
||||
"id": preset_id,
|
||||
"name": name,
|
||||
"sessionId": session_id,
|
||||
"trialId": trial_id,
|
||||
"rewardConfig": reward_config,
|
||||
"createdAt": at,
|
||||
}
|
||||
|
||||
def get_preset(self, preset_id: str) -> dict:
|
||||
row = self.connection().execute("SELECT * FROM presets WHERE id=?", (preset_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(preset_id)
|
||||
return {
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
|
||||
def list_presets(self) -> list[dict]:
|
||||
rows = (
|
||||
self.connection().execute("SELECT * FROM presets ORDER BY created_at DESC").fetchall()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Lazy Optuna study integration used for durable trial history and pruning metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
class OptunaStudies:
|
||||
def __init__(self, root: Path):
|
||||
self.path = (root / "optuna.sqlite3").resolve()
|
||||
self.url = f"sqlite:///{self.path}"
|
||||
|
||||
def _study(self, session_id: str):
|
||||
import optuna
|
||||
|
||||
optuna.logging.set_verbosity(optuna.logging.WARNING)
|
||||
return optuna.create_study(
|
||||
study_name=f"reward-tuning-{session_id}",
|
||||
storage=self.url,
|
||||
direction="maximize",
|
||||
load_if_exists=True,
|
||||
pruner=optuna.pruners.SuccessiveHalvingPruner(
|
||||
min_resource=300, reduction_factor=3, min_early_stopping_rate=0
|
||||
),
|
||||
)
|
||||
|
||||
def record(
|
||||
self,
|
||||
session_id: str,
|
||||
reward_config: dict[str, Any],
|
||||
score: float,
|
||||
eligible: bool,
|
||||
rung: int,
|
||||
) -> int:
|
||||
"""Record an externally proposed Agent config through Optuna ask/tell."""
|
||||
import optuna
|
||||
|
||||
study = self._study(session_id)
|
||||
trial = study.ask()
|
||||
trial.set_user_attr("reward_config", reward_config)
|
||||
trial.set_user_attr("eligible", eligible)
|
||||
trial.set_user_attr("rung", rung)
|
||||
state = optuna.trial.TrialState.COMPLETE if eligible else optuna.trial.TrialState.PRUNED
|
||||
study.tell(trial, score if eligible else None, state=state)
|
||||
return trial.number
|
||||
|
||||
def summary(self, session_id: str) -> dict[str, Any]:
|
||||
study = self._study(session_id)
|
||||
completed = [trial for trial in study.trials if trial.value is not None]
|
||||
return {
|
||||
"studyName": study.study_name,
|
||||
"trialCount": len(study.trials),
|
||||
"bestValue": max((trial.value for trial in completed), default=None),
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
"""TensorBoard scalar ingestion with optional dependency isolation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from .storage import TuningStorage
|
||||
|
||||
|
||||
class TensorboardUnavailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def ingest_scalars(storage: TuningStorage, trial_id: str, log_dir: Path) -> int:
|
||||
"""Reload all scalar events and idempotently upsert them into SQLite."""
|
||||
try:
|
||||
from tensorboard.backend.event_processing.event_accumulator import EventAccumulator
|
||||
except ImportError as error:
|
||||
raise TensorboardUnavailable(
|
||||
"缺少 tensorboard,请安装 training_server/requirements.txt"
|
||||
) from error
|
||||
if not log_dir.is_dir():
|
||||
return 0
|
||||
accumulator = EventAccumulator(str(log_dir), size_guidance={"scalars": 0})
|
||||
try:
|
||||
accumulator.Reload()
|
||||
except (OSError, ValueError):
|
||||
return 0
|
||||
points: list[tuple[str, int, float, float]] = []
|
||||
for tag in accumulator.Tags().get("scalars", []):
|
||||
for event in accumulator.Scalars(tag):
|
||||
points.append((tag, int(event.step), float(event.wall_time), float(event.value)))
|
||||
if points:
|
||||
storage.insert_metrics(trial_id, points)
|
||||
return len(points)
|
||||
Reference in New Issue
Block a user