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
+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)