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

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