feat(training): release V0.8 自调参 Agent
This commit is contained in:
@@ -0,0 +1,26 @@
|
||||
"""Reward auto-tuning support for the local training service."""
|
||||
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
PARAMETER_SPECS,
|
||||
WEIGHT_SPECS,
|
||||
RewardConfigError,
|
||||
apply_reward_configuration,
|
||||
merge_proposal,
|
||||
validate_configuration,
|
||||
validate_proposal,
|
||||
)
|
||||
from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation
|
||||
|
||||
__all__ = [
|
||||
"BASE_REWARD_CONFIGURATION",
|
||||
"DEFAULT_OBJECTIVE_WEIGHTS",
|
||||
"PARAMETER_SPECS",
|
||||
"WEIGHT_SPECS",
|
||||
"RewardConfigError",
|
||||
"apply_reward_configuration",
|
||||
"merge_proposal",
|
||||
"score_evaluation",
|
||||
"validate_configuration",
|
||||
"validate_proposal",
|
||||
]
|
||||
@@ -0,0 +1,134 @@
|
||||
"""PydanticAI adapter for the DeepSeek reward-tuning advisor."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from .schema import validate_proposal
|
||||
|
||||
SYSTEM_PROMPT = """你是 Unitree Go2 强化学习奖励调参专家。
|
||||
只根据提供的数值配置、训练曲线摘要和固定评估结果提出下一轮稀疏修改。
|
||||
必须优先保持速度跟踪与跌倒安全门槛;每轮最多修改四个白名单标量,不得改变符号、函数、传感器或结构。
|
||||
不要建议 Python 代码、命令、文件路径或白名单外参数。输出必须符合 RewardProposal schema。
|
||||
"""
|
||||
|
||||
|
||||
class AdvisorUnavailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AdvisorConfig:
|
||||
api_key: str | None
|
||||
base_url: str = "https://api.deepseek.com"
|
||||
model: str = "deepseek-v4-flash"
|
||||
|
||||
@classmethod
|
||||
def from_environment(cls) -> AdvisorConfig:
|
||||
return cls(
|
||||
api_key=os.environ.get("DEEPSEEK_API_KEY"),
|
||||
base_url=os.environ.get("MUJOCO_TUNING_AGENT_BASE_URL", "https://api.deepseek.com"),
|
||||
model=os.environ.get("MUJOCO_TUNING_AGENT_MODEL", "deepseek-v4-flash"),
|
||||
)
|
||||
|
||||
|
||||
class DeepSeekAdvisor:
|
||||
def __init__(self, config: AdvisorConfig | None = None):
|
||||
self.config = config or AdvisorConfig.from_environment()
|
||||
self._cached_agent = None
|
||||
|
||||
def capability(self) -> dict[str, Any]:
|
||||
try:
|
||||
import pydantic_ai # noqa: F401
|
||||
except ImportError:
|
||||
installed = False
|
||||
else:
|
||||
installed = True
|
||||
return {
|
||||
"configured": bool(self.config.api_key) and installed,
|
||||
"apiKeyConfigured": bool(self.config.api_key),
|
||||
"frameworkInstalled": installed,
|
||||
"model": self.config.model,
|
||||
"baseUrl": self.config.base_url,
|
||||
}
|
||||
|
||||
def _agent(self):
|
||||
if self._cached_agent is not None:
|
||||
return self._cached_agent
|
||||
if not self.config.api_key:
|
||||
raise AdvisorUnavailable("未配置 DEEPSEEK_API_KEY")
|
||||
try:
|
||||
import httpx2
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic_ai import Agent, PromptedOutput
|
||||
from pydantic_ai.models.openai import OpenAIChatModel
|
||||
from pydantic_ai.providers.openai import OpenAIProvider
|
||||
except ImportError as error:
|
||||
raise AdvisorUnavailable(
|
||||
"缺少 PydanticAI,请安装 training_server/requirements.txt"
|
||||
) from error
|
||||
|
||||
class RewardProposalOutput(BaseModel):
|
||||
weights: dict[str, float] = Field(default_factory=dict)
|
||||
params: dict[str, float] = Field(default_factory=dict)
|
||||
rationale: str = Field(min_length=1, max_length=2000)
|
||||
expected_impact: dict[str, str] = Field(default_factory=dict)
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
|
||||
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("ALL_PROXY")
|
||||
if proxy and proxy.startswith("socks://"):
|
||||
proxy = "socks5://" + proxy.removeprefix("socks://")
|
||||
http_client = httpx2.AsyncClient(proxy=proxy, trust_env=False, timeout=60.0)
|
||||
provider = OpenAIProvider(
|
||||
base_url=self.config.base_url, api_key=self.config.api_key, http_client=http_client
|
||||
)
|
||||
model = OpenAIChatModel(self.config.model, provider=provider) # type: ignore[arg-type]
|
||||
self._cached_agent = Agent(
|
||||
model,
|
||||
output_type=PromptedOutput(RewardProposalOutput),
|
||||
system_prompt=SYSTEM_PROMPT,
|
||||
retries=2,
|
||||
model_settings={"temperature": 0.2},
|
||||
)
|
||||
return self._cached_agent
|
||||
|
||||
def propose(self, context: dict[str, Any], previous: dict) -> dict[str, Any]:
|
||||
prompt = json.dumps(context, ensure_ascii=False, separators=(",", ":"), allow_nan=False)
|
||||
result = self._agent().run_sync(prompt)
|
||||
output = result.output
|
||||
patch = validate_proposal(
|
||||
{"weights": dict(output.weights), "params": dict(output.params)}, previous
|
||||
)
|
||||
try:
|
||||
usage = result.usage()
|
||||
usage_value = {
|
||||
key: getattr(usage, key)
|
||||
for key in ("requests", "input_tokens", "output_tokens", "total_tokens")
|
||||
if getattr(usage, key, None) is not None
|
||||
}
|
||||
except (AttributeError, TypeError):
|
||||
usage_value = {}
|
||||
return {
|
||||
"patch": patch,
|
||||
"rationale": output.rationale,
|
||||
"expectedImpact": dict(output.expected_impact),
|
||||
"confidence": float(output.confidence),
|
||||
"promptHash": hashlib.sha256(prompt.encode()).hexdigest(),
|
||||
"usage": usage_value,
|
||||
"model": self.config.model,
|
||||
}
|
||||
|
||||
def test_connection(self) -> dict[str, Any]:
|
||||
base = {
|
||||
"weights": {"track_linear_velocity": 1.0},
|
||||
"params": {},
|
||||
"instruction": "仅返回一个合法示例:把 track_linear_velocity 改为 1.1。",
|
||||
}
|
||||
# A minimal full previous config is supplied by callers for actual proposals;
|
||||
# connectivity probing only verifies the provider and structured response path.
|
||||
result = self._agent().run_sync(json.dumps(base, ensure_ascii=False))
|
||||
return {"ok": True, "model": self.config.model, "outputType": type(result.output).__name__}
|
||||
@@ -0,0 +1,715 @@
|
||||
"""Persistent reward tuning session orchestrator."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .advisor import AdvisorUnavailable, DeepSeekAdvisor
|
||||
from .process import GpuLease, ResourceBusyError, terminate_process
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
merge_proposal,
|
||||
validate_proposal,
|
||||
)
|
||||
from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation, validate_objective_weights
|
||||
from .storage import TuningStorage, now_iso
|
||||
from .study import OptunaStudies
|
||||
from .tensorboard import ingest_scalars
|
||||
|
||||
ACTIVE_SESSION_STATES = {
|
||||
"queued",
|
||||
"running",
|
||||
"evaluating",
|
||||
"awaiting_approval",
|
||||
"paused",
|
||||
"interrupted",
|
||||
}
|
||||
RUNNING_STATES = {"queued", "running", "evaluating"}
|
||||
RUN_NAME = re.compile(r"^[A-Za-z0-9_.-]{1,64}$")
|
||||
|
||||
|
||||
class TuningError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class TuningManager:
|
||||
def __init__(
|
||||
self,
|
||||
trainer_root: Path,
|
||||
python: str,
|
||||
data_root: Path,
|
||||
lease: GpuLease,
|
||||
advisor: Any | None = None,
|
||||
):
|
||||
self.trainer_root = trainer_root.expanduser().resolve()
|
||||
self.python = python
|
||||
self.data_root = data_root.expanduser().resolve()
|
||||
self.data_root.mkdir(parents=True, exist_ok=True)
|
||||
self.storage = TuningStorage(self.data_root / "tuning.sqlite3")
|
||||
self.storage.recover_interrupted()
|
||||
self.studies = OptunaStudies(self.data_root)
|
||||
self.lease = lease
|
||||
self.advisor = advisor or DeepSeekAdvisor()
|
||||
self.lock = threading.RLock()
|
||||
self.condition = threading.Condition(self.lock)
|
||||
self.workers: dict[str, threading.Thread] = {}
|
||||
self.processes: dict[str, subprocess.Popen[str]] = {}
|
||||
self.cancel_events: dict[str, threading.Event] = {}
|
||||
|
||||
def capability(self) -> dict[str, Any]:
|
||||
capability = self.advisor.capability()
|
||||
capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()})
|
||||
return capability
|
||||
|
||||
@staticmethod
|
||||
def _integer(payload: dict, name: str, default: int, minimum: int, maximum: int) -> int:
|
||||
value = payload.get(name, default)
|
||||
if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum:
|
||||
raise TuningError(f"{name} 必须在 {minimum}–{maximum} 之间")
|
||||
return value
|
||||
|
||||
def parse_create(self, payload: Any) -> tuple[str, dict, dict, bool]:
|
||||
if not isinstance(payload, dict):
|
||||
raise TuningError("请求体必须是 JSON 对象")
|
||||
mode = payload.get("mode", "automatic")
|
||||
if mode not in ("automatic", "approval"):
|
||||
raise TuningError("mode 必须是 automatic 或 approval")
|
||||
if payload.get("taskId", "Unitree-Go2-Flat") != "Unitree-Go2-Flat":
|
||||
raise TuningError("第一版只支持 Unitree-Go2-Flat")
|
||||
run_name = payload.get("runName", "auto-tune")
|
||||
if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name):
|
||||
raise TuningError("runName 格式无效")
|
||||
gpu_ids = payload.get("gpuIds", [0])
|
||||
if (
|
||||
not isinstance(gpu_ids, list)
|
||||
or not gpu_ids
|
||||
or any(
|
||||
isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 255
|
||||
for value in gpu_ids
|
||||
)
|
||||
):
|
||||
raise TuningError("gpuIds 必须是非空非负整数数组")
|
||||
trial_count = self._integer(payload, "trialCount", 12, 4, 20)
|
||||
rung0 = self._integer(payload, "initialIterations", 300, 1, 1000000)
|
||||
rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000)
|
||||
rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000)
|
||||
config = {
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384),
|
||||
"seed": self._integer(payload, "seed", 42, 0, 2147483647),
|
||||
"runName": run_name,
|
||||
"gpuIds": gpu_ids,
|
||||
"trialCount": trial_count,
|
||||
"rungs": [rung0, rung1, rung2],
|
||||
"promote": [trial_count, min(4, trial_count), min(2, trial_count)],
|
||||
"evalNumEnvs": self._integer(payload, "evalNumEnvs", 256, 1, 4096),
|
||||
"evalSteps": self._integer(payload, "evalSteps", 1000, 10, 100000),
|
||||
"earlyStopPatience": self._integer(payload, "earlyStopPatience", 4, 1, 20),
|
||||
}
|
||||
objective = validate_objective_weights(
|
||||
payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS)
|
||||
)
|
||||
fallback = payload.get("fallbackEnabled", False)
|
||||
if not isinstance(fallback, bool):
|
||||
raise TuningError("fallbackEnabled 必须是布尔值")
|
||||
return mode, config, objective, fallback
|
||||
|
||||
def create(self, payload: Any) -> dict:
|
||||
mode, config, objective, fallback = self.parse_create(payload)
|
||||
capability = self.capability()
|
||||
if not capability["ready"]:
|
||||
raise TuningError("评估入口未就绪")
|
||||
if not capability["configured"] and not fallback:
|
||||
raise TuningError("DeepSeek Agent 未配置;设置 DEEPSEEK_API_KEY 或显式启用 fallback")
|
||||
with self.lock:
|
||||
active = [
|
||||
session
|
||||
for session in self.storage.list_sessions()
|
||||
if session["state"] in ACTIVE_SESSION_STATES
|
||||
]
|
||||
if active:
|
||||
raise ResourceBusyError("已有调参 session 未结束")
|
||||
session = self.storage.create_session(mode, config, objective, fallback)
|
||||
self._start_worker(session["id"], resume=False)
|
||||
return self.detail(session["id"])
|
||||
|
||||
def _start_worker(self, session_id: str, resume: bool) -> None:
|
||||
cancel = threading.Event()
|
||||
self.cancel_events[session_id] = cancel
|
||||
worker = threading.Thread(
|
||||
target=self._run_session,
|
||||
args=(session_id, resume, cancel),
|
||||
name=f"tuning-{session_id[:8]}",
|
||||
daemon=True,
|
||||
)
|
||||
self.workers[session_id] = worker
|
||||
worker.start()
|
||||
|
||||
def detail(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
session["trials"] = self.storage.list_trials(session_id)
|
||||
session["proposals"] = self.storage.list_proposals(session_id)
|
||||
session["audit"] = self.storage.audit_events(session_id)
|
||||
return session
|
||||
|
||||
def list(self) -> list[dict]:
|
||||
return self.storage.list_sessions()
|
||||
|
||||
def _session_root(self, session_id: str) -> Path:
|
||||
root = (self.data_root / "sessions" / session_id).resolve()
|
||||
if not root.is_relative_to(self.data_root):
|
||||
raise TuningError("非法 session 路径")
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
|
||||
def _run_command(
|
||||
self,
|
||||
session_id: str,
|
||||
command: list[str],
|
||||
cwd: Path,
|
||||
environment: dict[str, str],
|
||||
log_path: Path,
|
||||
) -> int:
|
||||
owner = f"tuning:{session_id}"
|
||||
self.lease.acquire(owner)
|
||||
try:
|
||||
process = subprocess.Popen(
|
||||
command,
|
||||
cwd=cwd,
|
||||
env=environment,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
bufsize=1,
|
||||
start_new_session=True,
|
||||
)
|
||||
with self.lock:
|
||||
self.processes[session_id] = process
|
||||
assert process.stdout is not None
|
||||
with log_path.open("a", encoding="utf-8") as log:
|
||||
for line in process.stdout:
|
||||
log.write(line)
|
||||
log.flush()
|
||||
if self.cancel_events[session_id].is_set():
|
||||
terminate_process(process)
|
||||
break
|
||||
return process.wait()
|
||||
finally:
|
||||
with self.lock:
|
||||
self.processes.pop(session_id, None)
|
||||
self.lease.release(owner)
|
||||
|
||||
@staticmethod
|
||||
def _latest_checkpoint(run_dir: Path) -> Path | None:
|
||||
values = []
|
||||
for path in run_dir.glob("model_*.pt"):
|
||||
match = re.fullmatch(r"model_(\d+)\.pt", path.name)
|
||||
if match:
|
||||
values.append((int(match.group(1)), path))
|
||||
return max(values, default=(0, None), key=lambda value: value[0])[1]
|
||||
|
||||
def _execute_trial(
|
||||
self, session: dict, trial: dict, resume_checkpoint: Path | None = None
|
||||
) -> dict:
|
||||
session_id, trial_id = session["id"], trial["id"]
|
||||
root = self._session_root(session_id)
|
||||
run_dir = (root / trial["runDir"]).resolve()
|
||||
if not run_dir.is_relative_to(root):
|
||||
raise TuningError("trial 目录越界")
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
reward_path = run_dir / "reward_config.json"
|
||||
reward_path.write_text(
|
||||
json.dumps(trial["rewardConfig"], ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||
)
|
||||
config = session["config"]
|
||||
command = [
|
||||
self.python,
|
||||
"-u",
|
||||
"scripts/train.py",
|
||||
config["taskId"],
|
||||
f"--env.scene.num-envs={config['numEnvs']}",
|
||||
f"--agent.max-iterations={trial['targetIterations']}",
|
||||
f"--agent.seed={config['seed']}",
|
||||
f"--agent.run-name={config['runName']}-t{trial['number']}-r{trial['rung']}",
|
||||
"--agent.logger=tensorboard",
|
||||
"--agent.upload-model=False",
|
||||
"--gpu-ids",
|
||||
json.dumps(config["gpuIds"], separators=(",", ":")),
|
||||
"--output-dir",
|
||||
str(run_dir),
|
||||
"--reward-config",
|
||||
str(reward_path),
|
||||
]
|
||||
if resume_checkpoint is not None:
|
||||
command.extend(("--resume-checkpoint", str(resume_checkpoint)))
|
||||
environment = os.environ.copy()
|
||||
environment["WANDB_MODE"] = "disabled"
|
||||
environment["WANDB_SILENT"] = "true"
|
||||
self.storage.update_trial(
|
||||
trial_id, state="training", started_at=now_iso(), message="正在训练"
|
||||
)
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="running",
|
||||
message=f"正在训练 trial {trial['number']} / rung {trial['rung']}",
|
||||
)
|
||||
return_code = self._run_command(
|
||||
session_id, command, self.trainer_root, environment, run_dir / "train.log"
|
||||
)
|
||||
ingest_scalars(self.storage, trial_id, run_dir)
|
||||
if self.cancel_events[session_id].is_set():
|
||||
raise TuningError("session 已取消")
|
||||
checkpoint = self._latest_checkpoint(run_dir)
|
||||
policy = run_dir / "policy.onnx"
|
||||
if return_code != 0 or checkpoint is None or not policy.is_file():
|
||||
raise TuningError(f"训练失败(返回码 {return_code})或缺少 checkpoint/policy.onnx")
|
||||
self._wait_if_paused(session_id, self.cancel_events[session_id])
|
||||
|
||||
eval_output = run_dir / "evaluation.json"
|
||||
eval_command = [
|
||||
self.python,
|
||||
"-u",
|
||||
"scripts/evaluate.py",
|
||||
config["taskId"],
|
||||
"--checkpoint",
|
||||
str(checkpoint),
|
||||
"--output",
|
||||
str(eval_output),
|
||||
"--reward-config",
|
||||
str(reward_path),
|
||||
f"--num-envs={config['evalNumEnvs']}",
|
||||
f"--steps-per-seed={config['evalSteps']}",
|
||||
"--gpu-ids",
|
||||
json.dumps(config["gpuIds"], separators=(",", ":")),
|
||||
]
|
||||
self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估")
|
||||
self.storage.update_session(
|
||||
session_id, state="evaluating", message=f"正在评估 trial {trial['number']}"
|
||||
)
|
||||
return_code = self._run_command(
|
||||
session_id, eval_command, self.trainer_root, environment, run_dir / "evaluate.log"
|
||||
)
|
||||
ingest_scalars(self.storage, trial_id, run_dir / "evaluation-events")
|
||||
if return_code != 0 or not eval_output.is_file():
|
||||
raise TuningError(f"评估失败(返回码 {return_code})")
|
||||
evaluation = json.loads(eval_output.read_text(encoding="utf-8"))
|
||||
baseline_trial = self.storage.list_trials(session_id)[0]
|
||||
if baseline_trial["evaluation"] is None:
|
||||
scored = {
|
||||
"eligible": True,
|
||||
"score": 0.0,
|
||||
"components": {},
|
||||
"metrics": evaluation["metrics"],
|
||||
}
|
||||
else:
|
||||
scored = score_evaluation(
|
||||
baseline_trial["evaluation"]["metrics"],
|
||||
evaluation["metrics"],
|
||||
session["objectiveWeights"],
|
||||
)
|
||||
evaluation["score"] = scored
|
||||
rel_checkpoint = str(checkpoint.relative_to(root))
|
||||
rel_policy = str(policy.relative_to(root))
|
||||
self.storage.update_trial(
|
||||
trial_id,
|
||||
state="completed",
|
||||
ended_at=now_iso(),
|
||||
message="训练与评估完成",
|
||||
checkpoint_path=rel_checkpoint,
|
||||
policy_path=rel_policy,
|
||||
evaluation=evaluation,
|
||||
score=scored["score"],
|
||||
eligible=scored["eligible"],
|
||||
)
|
||||
try:
|
||||
optuna_number = self.studies.record(
|
||||
session_id,
|
||||
trial["rewardConfig"],
|
||||
scored["score"],
|
||||
scored["eligible"],
|
||||
trial["rung"],
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id, "optuna_trial_recorded", {"trialId": trial_id, "number": optuna_number}
|
||||
)
|
||||
except Exception as error:
|
||||
self.storage.audit(
|
||||
session_id, "optuna_record_failed", {"trialId": trial_id, "error": str(error)[:500]}
|
||||
)
|
||||
return self.storage.get_trial(trial_id)
|
||||
|
||||
def _best(self, session_id: str, rung: int | None = None) -> dict | None:
|
||||
trials = [
|
||||
trial
|
||||
for trial in self.storage.list_trials(session_id)
|
||||
if trial["state"] == "completed" and trial["eligible"]
|
||||
]
|
||||
if rung is not None:
|
||||
trials = [trial for trial in trials if trial["rung"] == rung]
|
||||
return max(
|
||||
trials,
|
||||
key=lambda trial: trial["score"] if trial["score"] is not None else -999,
|
||||
default=None,
|
||||
)
|
||||
|
||||
def _proposal_context(self, session: dict) -> dict:
|
||||
trials = self.storage.list_trials(session["id"])[-12:]
|
||||
rejected_feedback = [
|
||||
proposal["feedback"]
|
||||
for proposal in self.storage.list_proposals(session["id"])
|
||||
if proposal["state"] == "rejected" and proposal["feedback"]
|
||||
][-4:]
|
||||
return {
|
||||
"task": session["config"]["taskId"],
|
||||
"objectiveWeights": session["objectiveWeights"],
|
||||
"allowlist": "服务端将验证固定 schema;最多四项修改",
|
||||
"rejectedFeedback": rejected_feedback,
|
||||
"trials": [
|
||||
{
|
||||
"number": t["number"],
|
||||
"rung": t["rung"],
|
||||
"score": t["score"],
|
||||
"eligible": t["eligible"],
|
||||
"rewardConfig": t["rewardConfig"],
|
||||
"evaluation": t["evaluation"] and t["evaluation"].get("metrics"),
|
||||
}
|
||||
for t in trials
|
||||
],
|
||||
}
|
||||
|
||||
def _fallback_patch(self, previous: dict, index: int) -> dict:
|
||||
names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip")
|
||||
name = names[index % len(names)]
|
||||
old = previous["weights"][name]
|
||||
factor = 1.1 if index % 2 == 0 else 0.9
|
||||
return validate_proposal({"weights": {name: old * factor}}, previous)
|
||||
|
||||
def _request_proposal(
|
||||
self, session: dict, previous: dict, base_trial_id: str, index: int
|
||||
) -> dict:
|
||||
try:
|
||||
result = self.advisor.propose(self._proposal_context(session), previous)
|
||||
source = "agent"
|
||||
except Exception as error:
|
||||
if not session["fallbackEnabled"]:
|
||||
raise AdvisorUnavailable(str(error)) from error
|
||||
result = {
|
||||
"patch": self._fallback_patch(previous, index),
|
||||
"rationale": f"Agent 不可用,显式 fallback:{error}",
|
||||
"expectedImpact": {},
|
||||
"confidence": 0.2,
|
||||
"promptHash": None,
|
||||
"usage": {},
|
||||
"model": "optuna-fallback",
|
||||
}
|
||||
source = "fallback"
|
||||
proposal = self.storage.create_proposal(
|
||||
session["id"],
|
||||
base_trial_id,
|
||||
result["patch"],
|
||||
result["rationale"],
|
||||
result.get("expectedImpact", {}),
|
||||
result["confidence"],
|
||||
source,
|
||||
)
|
||||
self.storage.audit(
|
||||
session["id"],
|
||||
"proposal_created",
|
||||
{
|
||||
"proposalId": proposal["id"],
|
||||
"source": source,
|
||||
"promptHash": result.get("promptHash"),
|
||||
"usage": result.get("usage", {}),
|
||||
"model": result.get("model"),
|
||||
},
|
||||
)
|
||||
return proposal
|
||||
|
||||
def _wait_for_approval(
|
||||
self, session_id: str, proposal_id: str, cancel: threading.Event
|
||||
) -> dict:
|
||||
with self.condition:
|
||||
while not cancel.is_set():
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["state"] != "pending":
|
||||
return proposal
|
||||
self.condition.wait(timeout=1.0)
|
||||
raise TuningError("session 已取消")
|
||||
|
||||
def _run_session(self, session_id: str, resume: bool, cancel: threading.Event) -> None:
|
||||
try:
|
||||
session = self.storage.get_session(session_id)
|
||||
trials = self.storage.list_trials(session_id)
|
||||
if resume:
|
||||
root = self._session_root(session_id)
|
||||
for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]:
|
||||
run_dir = (root / interrupted["runDir"]).resolve()
|
||||
if run_dir.is_relative_to(root):
|
||||
shutil.rmtree(run_dir, ignore_errors=True)
|
||||
self.storage.delete_trial(interrupted["id"])
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"session_resumed",
|
||||
{
|
||||
"discardedInterruptedTrials": [
|
||||
t["id"] for t in trials if t["state"] == "interrupted"
|
||||
]
|
||||
},
|
||||
)
|
||||
trials = self.storage.list_trials(session_id)
|
||||
if not trials:
|
||||
baseline_dir = "trial-000-rung-0"
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
0,
|
||||
0,
|
||||
session["config"]["rungs"][0],
|
||||
deepcopy(BASE_REWARD_CONFIGURATION),
|
||||
None,
|
||||
baseline_dir,
|
||||
)
|
||||
self._execute_trial(session, trial)
|
||||
session = self.storage.get_session(session_id)
|
||||
completed_rung0 = [
|
||||
t
|
||||
for t in self.storage.list_trials(session_id)
|
||||
if t["rung"] == 0 and t["state"] == "completed"
|
||||
]
|
||||
next_number = len({t["number"] for t in completed_rung0})
|
||||
best_score = max((t["score"] or 0.0 for t in completed_rung0), default=0.0)
|
||||
no_improve = 0
|
||||
while (
|
||||
next_number < session["config"]["trialCount"]
|
||||
and no_improve < session["config"]["earlyStopPatience"]
|
||||
):
|
||||
if cancel.is_set():
|
||||
raise TuningError("session 已取消")
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
base = self._best(session_id, rung=0) or completed_rung0[0]
|
||||
proposal = self._request_proposal(
|
||||
session, base["rewardConfig"], base["id"], next_number
|
||||
)
|
||||
if session["mode"] == "approval":
|
||||
self.storage.update_session(
|
||||
session_id, state="awaiting_approval", message="等待批准 Agent 建议"
|
||||
)
|
||||
proposal = self._wait_for_approval(session_id, proposal["id"], cancel)
|
||||
if proposal["state"] == "rejected":
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"proposal_rejected",
|
||||
{"proposalId": proposal["id"], "feedback": proposal["feedback"]},
|
||||
)
|
||||
continue
|
||||
else:
|
||||
self.storage.decide_proposal(proposal["id"], "approved", "自动模式")
|
||||
proposal = self.storage.get_proposal(proposal["id"])
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
reward_config = merge_proposal(base["rewardConfig"], proposal["patch"])
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
next_number,
|
||||
0,
|
||||
session["config"]["rungs"][0],
|
||||
reward_config,
|
||||
proposal["id"],
|
||||
f"trial-{next_number:03d}-rung-0",
|
||||
)
|
||||
result = self._execute_trial(session, trial)
|
||||
if result["eligible"] and (result["score"] or -999) > best_score + 0.01:
|
||||
best_score = result["score"]
|
||||
no_improve = 0
|
||||
else:
|
||||
no_improve += 1
|
||||
self.storage.update_session(session_id, consecutive_no_improve=no_improve)
|
||||
next_number += 1
|
||||
|
||||
# Promote top configurations; each new rung resumes its own previous checkpoint.
|
||||
for rung in (1, 2):
|
||||
self._wait_if_paused(session_id, cancel)
|
||||
previous = [
|
||||
t
|
||||
for t in self.storage.list_trials(session_id)
|
||||
if t["rung"] == rung - 1 and t["state"] == "completed" and t["eligible"]
|
||||
]
|
||||
previous.sort(key=lambda t: t["score"] or -999, reverse=True)
|
||||
promoted_numbers = {
|
||||
trial["number"]
|
||||
for trial in self.storage.list_trials(session_id)
|
||||
if trial["rung"] == rung and trial["state"] == "completed"
|
||||
}
|
||||
for parent in previous[: session["config"]["promote"][rung]]:
|
||||
if parent["number"] in promoted_numbers:
|
||||
continue
|
||||
if cancel.is_set():
|
||||
raise TuningError("session 已取消")
|
||||
root = self._session_root(session_id)
|
||||
checkpoint = root / parent["checkpointPath"]
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
parent["number"],
|
||||
rung,
|
||||
session["config"]["rungs"][rung],
|
||||
parent["rewardConfig"],
|
||||
parent["proposalId"],
|
||||
f"trial-{parent['number']:03d}-rung-{rung}",
|
||||
)
|
||||
self._execute_trial(session, trial, checkpoint)
|
||||
|
||||
best = (
|
||||
self._best(session_id, rung=2)
|
||||
or self._best(session_id, rung=1)
|
||||
or self._best(session_id, rung=0)
|
||||
)
|
||||
if best is None:
|
||||
raise TuningError("没有通过安全门槛的 trial")
|
||||
preset_name = f"{session['config']['runName']}-{session_id[:8]}"
|
||||
self.storage.save_preset(preset_name, session_id, best["id"], best["rewardConfig"])
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="succeeded",
|
||||
best_trial_id=best["id"],
|
||||
current_trial_id=None,
|
||||
message="调参完成",
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id, "session_completed", {"bestTrialId": best["id"], "preset": preset_name}
|
||||
)
|
||||
except Exception as error:
|
||||
state = self.storage.get_session(session_id)["state"]
|
||||
if cancel.is_set() or state == "cancelled":
|
||||
self.storage.update_session(
|
||||
session_id, state="cancelled", message="调参已取消", current_trial_id=None
|
||||
)
|
||||
else:
|
||||
self.storage.update_session(
|
||||
session_id, state="failed", message=str(error), current_trial_id=None
|
||||
)
|
||||
self.storage.audit(session_id, "session_failed", {"error": str(error)[:1000]})
|
||||
finally:
|
||||
with self.lock:
|
||||
self.workers.pop(session_id, None)
|
||||
self.processes.pop(session_id, None)
|
||||
|
||||
def _wait_if_paused(self, session_id: str, cancel: threading.Event) -> None:
|
||||
with self.condition:
|
||||
while self.storage.get_session(session_id)["state"] == "paused" and not cancel.is_set():
|
||||
self.condition.wait(timeout=1.0)
|
||||
|
||||
def approve(self, session_id: str, proposal_id: str, payload: Any) -> dict:
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["sessionId"] != session_id:
|
||||
raise TuningError("proposal 不属于该 session")
|
||||
patch = proposal["patch"]
|
||||
feedback = None
|
||||
if isinstance(payload, dict):
|
||||
feedback = payload.get("feedback")
|
||||
if "patch" in payload:
|
||||
base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
patch = validate_proposal(payload["patch"], base["rewardConfig"])
|
||||
if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch):
|
||||
raise TuningError("proposal 已处理")
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"proposal_approved",
|
||||
{"proposalId": proposal_id, "modified": patch != proposal["patch"]},
|
||||
)
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def reject(self, session_id: str, proposal_id: str, payload: Any) -> dict:
|
||||
feedback = payload.get("feedback", "") if isinstance(payload, dict) else ""
|
||||
if not isinstance(feedback, str) or len(feedback) > 2000:
|
||||
raise TuningError("feedback 无效")
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
if proposal["sessionId"] != session_id:
|
||||
raise TuningError("proposal 不属于该 session")
|
||||
if not self.storage.decide_proposal(proposal_id, "rejected", feedback):
|
||||
raise TuningError("proposal 已处理")
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def pause(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in RUNNING_STATES | {"awaiting_approval"}:
|
||||
raise TuningError("当前状态不能暂停")
|
||||
self.storage.update_session(
|
||||
session_id, state="paused", message="已暂停后续调度;当前子进程将完成"
|
||||
)
|
||||
return self.detail(session_id)
|
||||
|
||||
def resume(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] == "paused":
|
||||
self.storage.update_session(session_id, state="running", message="继续调参")
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
elif session["state"] == "interrupted":
|
||||
self.storage.update_session(session_id, state="queued", message="从最近完整结果恢复")
|
||||
self._start_worker(session_id, resume=True)
|
||||
else:
|
||||
raise TuningError("当前状态不能恢复")
|
||||
return self.detail(session_id)
|
||||
|
||||
def cancel(self, session_id: str) -> dict:
|
||||
self.storage.get_session(session_id)
|
||||
self.storage.update_session(session_id, state="cancelled", message="正在取消")
|
||||
event = self.cancel_events.get(session_id)
|
||||
if event:
|
||||
event.set()
|
||||
process = self.processes.get(session_id)
|
||||
if process:
|
||||
terminate_process(process)
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def metrics(
|
||||
self, session_id: str, trial_id: str, tags: list[str] | None, max_points: int
|
||||
) -> dict:
|
||||
trial = self.storage.get_trial(trial_id)
|
||||
if trial["sessionId"] != session_id:
|
||||
raise TuningError("trial 不属于该 session")
|
||||
return {"trialId": trial_id, "series": self.storage.metrics(trial_id, tags, max_points)}
|
||||
|
||||
def best_artifact(self, session_id: str) -> Path:
|
||||
session = self.storage.get_session(session_id)
|
||||
if not session["bestTrialId"]:
|
||||
raise TuningError("尚无最佳策略")
|
||||
trial = self.storage.get_trial(session["bestTrialId"])
|
||||
if not trial["policyPath"]:
|
||||
raise TuningError("最佳策略文件不存在")
|
||||
root = self._session_root(session_id)
|
||||
path = (root / trial["policyPath"]).resolve()
|
||||
if not path.is_relative_to(root) or not path.is_file():
|
||||
raise TuningError("最佳策略文件不存在")
|
||||
return path
|
||||
|
||||
def preset_config(self, preset_id: str) -> dict:
|
||||
return self.storage.get_preset(preset_id)["rewardConfig"]
|
||||
|
||||
def test_agent(self) -> dict:
|
||||
try:
|
||||
return self.advisor.test_connection()
|
||||
except Exception as error:
|
||||
raise TuningError(f"Agent 连接测试失败:{error}") from error
|
||||
|
||||
def shutdown(self) -> None:
|
||||
for session_id in list(self.workers):
|
||||
with suppress(KeyError, TuningError):
|
||||
self.cancel(session_id)
|
||||
for worker in list(self.workers.values()):
|
||||
worker.join(timeout=7)
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Shared GPU lease and process-group lifecycle helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
|
||||
|
||||
class ResourceBusyError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class GpuLease:
|
||||
def __init__(self):
|
||||
self.lock = threading.RLock()
|
||||
self.owner: str | None = None
|
||||
|
||||
def acquire(self, owner: str) -> None:
|
||||
with self.lock:
|
||||
if self.owner is not None and self.owner != owner:
|
||||
raise ResourceBusyError(f"计算资源正由 {self.owner} 使用")
|
||||
self.owner = owner
|
||||
|
||||
def release(self, owner: str) -> None:
|
||||
with self.lock:
|
||||
if self.owner == owner:
|
||||
self.owner = None
|
||||
|
||||
def public(self) -> str | None:
|
||||
with self.lock:
|
||||
return self.owner
|
||||
|
||||
|
||||
def terminate_process(process: subprocess.Popen[str], grace_seconds: float = 5.0) -> None:
|
||||
if process.poll() is not None:
|
||||
return
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGTERM)
|
||||
try:
|
||||
process.wait(timeout=grace_seconds)
|
||||
except subprocess.TimeoutExpired:
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGKILL)
|
||||
@@ -0,0 +1,213 @@
|
||||
"""Pure-Python reward tuning schema shared by the service and trainer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
MAX_PROPOSAL_CHANGES = 4
|
||||
MIN_CHANGE_RATIO = 0.5
|
||||
MAX_CHANGE_RATIO = 2.0
|
||||
|
||||
|
||||
class RewardConfigError(ValueError):
|
||||
"""A reward configuration or proposal violated the allowlist."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NumericSpec:
|
||||
minimum: float
|
||||
maximum: float
|
||||
default: float
|
||||
allow_zero: bool = True
|
||||
|
||||
|
||||
WEIGHT_SPECS: dict[str, NumericSpec] = {
|
||||
"track_linear_velocity": NumericSpec(0.5, 3.0, 1.0, False),
|
||||
"track_angular_velocity": NumericSpec(0.25, 2.0, 1.0, False),
|
||||
"body_orientation_l2": NumericSpec(-3.0, -0.1, -1.0, False),
|
||||
"pose": NumericSpec(0.0, 2.5, 1.0),
|
||||
"body_ang_vel": NumericSpec(-0.2, 0.0, -0.05),
|
||||
"angular_momentum": NumericSpec(-0.1, 0.0, -0.025),
|
||||
"is_terminated": NumericSpec(-400.0, -50.0, -200.0, False),
|
||||
"joint_acc_l2": NumericSpec(-2.0e-6, 0.0, -2.5e-7),
|
||||
"joint_pos_limits": NumericSpec(-30.0, -2.0, -10.0, False),
|
||||
"action_rate_l2": NumericSpec(-0.2, -0.005, -0.05, False),
|
||||
"foot_gait": NumericSpec(0.0, 1.5, 0.5),
|
||||
"foot_clearance": NumericSpec(-3.0, 0.0, -1.0),
|
||||
"foot_slip": NumericSpec(-1.0, 0.0, -0.25),
|
||||
"soft_landing": NumericSpec(-5.0e-3, 0.0, -1.0e-3),
|
||||
"stand_still": NumericSpec(-3.0, 0.0, -1.0),
|
||||
"electrical_power": NumericSpec(-5.0e-3, 0.0, 0.0),
|
||||
}
|
||||
|
||||
PARAMETER_SPECS: dict[str, NumericSpec] = {
|
||||
"track_linear_velocity.std": NumericSpec(0.25, 1.0, math.sqrt(0.25), False),
|
||||
"track_angular_velocity.std": NumericSpec(0.35, 1.2, math.sqrt(0.5), False),
|
||||
"pose.std_standing_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.std_walking_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.std_running_scale": NumericSpec(0.5, 2.0, 1.0, False),
|
||||
"pose.walking_threshold": NumericSpec(0.05, 0.5, 0.1, False),
|
||||
"pose.running_threshold": NumericSpec(1.0, 2.5, 1.5, False),
|
||||
"foot_gait.period": NumericSpec(0.4, 0.8, 0.6, False),
|
||||
"foot_gait.threshold": NumericSpec(0.45, 0.65, 0.56, False),
|
||||
"foot_gait.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"foot_clearance.target_height": NumericSpec(0.06, 0.16, 0.1, False),
|
||||
"foot_clearance.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"foot_slip.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"soft_landing.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
"stand_still.command_threshold": NumericSpec(0.02, 0.3, 0.1, False),
|
||||
}
|
||||
|
||||
BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = {
|
||||
"weights": {name: spec.default for name, spec in WEIGHT_SPECS.items()},
|
||||
"params": {name: spec.default for name, spec in PARAMETER_SPECS.items()},
|
||||
}
|
||||
|
||||
|
||||
def _number(name: str, value: Any, spec: NumericSpec) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise RewardConfigError(f"{name} 必须是数值")
|
||||
result = float(value)
|
||||
if not math.isfinite(result):
|
||||
raise RewardConfigError(f"{name} 必须是有限数值")
|
||||
if result == 0.0 and not spec.allow_zero:
|
||||
raise RewardConfigError(f"{name} 不允许关闭")
|
||||
if result < spec.minimum or result > spec.maximum:
|
||||
raise RewardConfigError(f"{name} 必须在 {spec.minimum}–{spec.maximum} 之间")
|
||||
return result
|
||||
|
||||
|
||||
def _mapping(value: Any, name: str) -> Mapping[str, Any]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise RewardConfigError(f"{name} 必须是对象")
|
||||
return value
|
||||
|
||||
|
||||
def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None:
|
||||
params = config["params"]
|
||||
if params["pose.walking_threshold"] >= params["pose.running_threshold"]:
|
||||
raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold")
|
||||
|
||||
|
||||
def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a complete configuration and reject missing/unknown fields."""
|
||||
root = _mapping(value, "rewardConfig")
|
||||
if set(root) != {"weights", "params"}:
|
||||
raise RewardConfigError("rewardConfig 只能包含 weights 和 params")
|
||||
raw_weights = _mapping(root["weights"], "weights")
|
||||
raw_params = _mapping(root["params"], "params")
|
||||
if set(raw_weights) != set(WEIGHT_SPECS):
|
||||
raise RewardConfigError("weights 必须完整且不能包含未知奖励项")
|
||||
if set(raw_params) != set(PARAMETER_SPECS):
|
||||
raise RewardConfigError("params 必须完整且不能包含未知参数")
|
||||
config = {
|
||||
"weights": {
|
||||
name: _number(f"weights.{name}", raw_weights[name], spec)
|
||||
for name, spec in WEIGHT_SPECS.items()
|
||||
},
|
||||
"params": {
|
||||
name: _number(f"params.{name}", raw_params[name], spec)
|
||||
for name, spec in PARAMETER_SPECS.items()
|
||||
},
|
||||
}
|
||||
_cross_validate(config)
|
||||
return config
|
||||
|
||||
|
||||
def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config."""
|
||||
current = validate_configuration(previous)
|
||||
root = _mapping(value, "proposal")
|
||||
if not set(root).issubset({"weights", "params"}):
|
||||
raise RewardConfigError("proposal 只能包含 weights 和 params")
|
||||
raw_weights = _mapping(root.get("weights", {}), "weights")
|
||||
raw_params = _mapping(root.get("params", {}), "params")
|
||||
if len(raw_weights) + len(raw_params) == 0:
|
||||
raise RewardConfigError("proposal 至少需要一项修改")
|
||||
if len(raw_weights) + len(raw_params) > MAX_PROPOSAL_CHANGES:
|
||||
raise RewardConfigError(f"proposal 每轮最多修改 {MAX_PROPOSAL_CHANGES} 项")
|
||||
unknown_weights = set(raw_weights) - set(WEIGHT_SPECS)
|
||||
unknown_params = set(raw_params) - set(PARAMETER_SPECS)
|
||||
if unknown_weights:
|
||||
raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}")
|
||||
if unknown_params:
|
||||
raise RewardConfigError(f"未知奖励参数:{', '.join(sorted(unknown_params))}")
|
||||
|
||||
patch: dict[str, dict[str, float]] = {"weights": {}, "params": {}}
|
||||
for name, raw in raw_weights.items():
|
||||
value_number = _number(f"weights.{name}", raw, WEIGHT_SPECS[name])
|
||||
old = current["weights"][name]
|
||||
if old != 0.0 and value_number != 0.0:
|
||||
ratio = abs(value_number / old)
|
||||
if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO:
|
||||
raise RewardConfigError(
|
||||
f"weights.{name} 单轮变化必须在旧值幅度的 "
|
||||
f"{MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×"
|
||||
)
|
||||
if value_number == old:
|
||||
raise RewardConfigError(f"weights.{name} 没有发生变化")
|
||||
patch["weights"][name] = value_number
|
||||
for name, raw in raw_params.items():
|
||||
value_number = _number(f"params.{name}", raw, PARAMETER_SPECS[name])
|
||||
old = current["params"][name]
|
||||
ratio = abs(value_number / old)
|
||||
if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO:
|
||||
raise RewardConfigError(
|
||||
f"params.{name} 单轮变化必须在旧值的 {MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×"
|
||||
)
|
||||
if value_number == old:
|
||||
raise RewardConfigError(f"params.{name} 没有发生变化")
|
||||
patch["params"][name] = value_number
|
||||
|
||||
candidate = deepcopy(current)
|
||||
candidate["weights"].update(patch["weights"])
|
||||
candidate["params"].update(patch["params"])
|
||||
_cross_validate(candidate)
|
||||
return patch
|
||||
|
||||
|
||||
def merge_proposal(previous: Any, proposal: Any) -> dict[str, dict[str, float]]:
|
||||
current = validate_configuration(previous)
|
||||
patch = validate_proposal(proposal, current)
|
||||
merged = deepcopy(current)
|
||||
merged["weights"].update(patch["weights"])
|
||||
merged["params"].update(patch["params"])
|
||||
return validate_configuration(merged)
|
||||
|
||||
|
||||
def apply_reward_configuration(env_cfg: Any, value: Any) -> None:
|
||||
"""Apply a validated full config to a fresh mjlab environment config."""
|
||||
config = validate_configuration(value)
|
||||
for name, weight in config["weights"].items():
|
||||
if name not in env_cfg.rewards:
|
||||
raise RewardConfigError(f"环境缺少奖励项:{name}")
|
||||
env_cfg.rewards[name].weight = weight
|
||||
|
||||
params = config["params"]
|
||||
direct = {
|
||||
"track_linear_velocity.std": ("track_linear_velocity", "std"),
|
||||
"track_angular_velocity.std": ("track_angular_velocity", "std"),
|
||||
"pose.walking_threshold": ("pose", "walking_threshold"),
|
||||
"pose.running_threshold": ("pose", "running_threshold"),
|
||||
"foot_gait.period": ("foot_gait", "period"),
|
||||
"foot_gait.threshold": ("foot_gait", "threshold"),
|
||||
"foot_gait.command_threshold": ("foot_gait", "command_threshold"),
|
||||
"foot_clearance.target_height": ("foot_clearance", "target_height"),
|
||||
"foot_clearance.command_threshold": ("foot_clearance", "command_threshold"),
|
||||
"foot_slip.command_threshold": ("foot_slip", "command_threshold"),
|
||||
"soft_landing.command_threshold": ("soft_landing", "command_threshold"),
|
||||
"stand_still.command_threshold": ("stand_still", "command_threshold"),
|
||||
}
|
||||
for path, (term, parameter) in direct.items():
|
||||
env_cfg.rewards[term].params[parameter] = params[path]
|
||||
for regime in ("standing", "walking", "running"):
|
||||
key = f"std_{regime}"
|
||||
scale = params[f"pose.{key}_scale"]
|
||||
baseline = env_cfg.rewards["pose"].params[key]
|
||||
env_cfg.rewards["pose"].params[key] = {
|
||||
pattern: float(std) * scale for pattern, std in baseline.items()
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
"""Stable, reward-weight-independent evaluation scoring."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
DEFAULT_OBJECTIVE_WEIGHTS = {
|
||||
"velocity_tracking": 0.35,
|
||||
"action_smoothness": 0.20,
|
||||
"posture_stability": 0.15,
|
||||
"fall_avoidance": 0.15,
|
||||
"foot_slip": 0.10,
|
||||
"energy": 0.05,
|
||||
}
|
||||
|
||||
REQUIRED_METRICS = {
|
||||
"linear_velocity_rmse",
|
||||
"angular_velocity_rmse",
|
||||
"mean_action_acc",
|
||||
"orientation_error",
|
||||
"fall_rate",
|
||||
"slip_velocity",
|
||||
"mechanical_power",
|
||||
}
|
||||
|
||||
PHYSICAL_FLOORS = {
|
||||
"linear_velocity_rmse": 0.10,
|
||||
"angular_velocity_rmse": 0.10,
|
||||
"mean_action_acc": 0.01,
|
||||
"orientation_error": 0.05,
|
||||
"fall_rate": 0.02,
|
||||
"slip_velocity": 0.05,
|
||||
"mechanical_power": 10.0,
|
||||
}
|
||||
|
||||
|
||||
class EvaluationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def validate_objective_weights(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, Mapping) or set(value) != set(DEFAULT_OBJECTIVE_WEIGHTS):
|
||||
raise EvaluationError("objectiveWeights 必须完整包含六个目标")
|
||||
result: dict[str, float] = {}
|
||||
for key in DEFAULT_OBJECTIVE_WEIGHTS:
|
||||
raw = value[key]
|
||||
if isinstance(raw, bool) or not isinstance(raw, (int, float)):
|
||||
raise EvaluationError(f"objectiveWeights.{key} 必须是数值")
|
||||
number = float(raw)
|
||||
if not math.isfinite(number) or number < 0.0 or number > 1.0:
|
||||
raise EvaluationError(f"objectiveWeights.{key} 必须在 0–1 之间")
|
||||
result[key] = number
|
||||
if not math.isclose(sum(result.values()), 1.0, abs_tol=1.0e-6):
|
||||
raise EvaluationError("objectiveWeights 总和必须为 1")
|
||||
return result
|
||||
|
||||
|
||||
def validate_metrics(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise EvaluationError("metrics 必须是对象")
|
||||
missing = REQUIRED_METRICS - set(value)
|
||||
if missing:
|
||||
raise EvaluationError(f"metrics 缺少:{', '.join(sorted(missing))}")
|
||||
result: dict[str, float] = {}
|
||||
for key in REQUIRED_METRICS:
|
||||
raw = value[key]
|
||||
if isinstance(raw, bool) or not isinstance(raw, (int, float)):
|
||||
raise EvaluationError(f"metrics.{key} 必须是数值")
|
||||
number = float(raw)
|
||||
if not math.isfinite(number) or number < 0.0:
|
||||
raise EvaluationError(f"metrics.{key} 必须是非负有限数值")
|
||||
result[key] = number
|
||||
if result["fall_rate"] > 1.0:
|
||||
raise EvaluationError("metrics.fall_rate 必须在 0–1 之间")
|
||||
return result
|
||||
|
||||
|
||||
def _improvement(baseline: Mapping[str, float], current: Mapping[str, float], key: str) -> float:
|
||||
scale = max(abs(baseline[key]), PHYSICAL_FLOORS[key])
|
||||
return max(-1.0, min(1.0, (baseline[key] - current[key]) / scale))
|
||||
|
||||
|
||||
def score_evaluation(
|
||||
baseline_value: Any,
|
||||
current_value: Any,
|
||||
objective_weights: Any = DEFAULT_OBJECTIVE_WEIGHTS,
|
||||
) -> dict[str, Any]:
|
||||
baseline = validate_metrics(baseline_value)
|
||||
current = validate_metrics(current_value)
|
||||
weights = validate_objective_weights(objective_weights)
|
||||
components = {
|
||||
"velocity_tracking": 0.8 * _improvement(baseline, current, "linear_velocity_rmse")
|
||||
+ 0.2 * _improvement(baseline, current, "angular_velocity_rmse"),
|
||||
"action_smoothness": _improvement(baseline, current, "mean_action_acc"),
|
||||
"posture_stability": _improvement(baseline, current, "orientation_error"),
|
||||
"fall_avoidance": _improvement(baseline, current, "fall_rate"),
|
||||
"foot_slip": _improvement(baseline, current, "slip_velocity"),
|
||||
"energy": _improvement(baseline, current, "mechanical_power"),
|
||||
}
|
||||
tracking_limit = max(
|
||||
baseline["linear_velocity_rmse"] * 1.05, baseline["linear_velocity_rmse"] + 1.0e-6
|
||||
)
|
||||
eligible = (
|
||||
current["fall_rate"] <= baseline["fall_rate"] + 0.02
|
||||
and current["linear_velocity_rmse"] <= tracking_limit
|
||||
)
|
||||
total = sum(weights[key] * components[key] for key in weights)
|
||||
if not eligible:
|
||||
total = min(total, -1.0)
|
||||
return {
|
||||
"eligible": eligible,
|
||||
"score": total,
|
||||
"components": components,
|
||||
"metrics": current,
|
||||
"baselineMetrics": baseline,
|
||||
}
|
||||
@@ -0,0 +1,521 @@
|
||||
"""SQLite persistence for tuning sessions, trials, proposals and scalar data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
|
||||
|
||||
def now_iso() -> str:
|
||||
return datetime.now(UTC).isoformat()
|
||||
|
||||
|
||||
def _json(value: Any) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False)
|
||||
|
||||
|
||||
def _decode(value: str | None) -> Any:
|
||||
return json.loads(value) if value else None
|
||||
|
||||
|
||||
def _lttb(points: list[dict], threshold: int) -> list[dict]:
|
||||
"""Largest-Triangle-Three-Buckets downsampling preserving peaks and endpoints."""
|
||||
if threshold >= len(points) or threshold < 3:
|
||||
return points[:threshold]
|
||||
sampled = [points[0]]
|
||||
bucket_width = (len(points) - 2) / (threshold - 2)
|
||||
anchor_index = 0
|
||||
for bucket in range(threshold - 2):
|
||||
average_start = int((bucket + 1) * bucket_width) + 1
|
||||
average_end = min(int((bucket + 2) * bucket_width) + 1, len(points))
|
||||
average_bucket = points[average_start:average_end] or [points[-1]]
|
||||
average_x = sum(point["step"] for point in average_bucket) / len(average_bucket)
|
||||
average_y = sum(point["value"] for point in average_bucket) / len(average_bucket)
|
||||
range_start = int(bucket * bucket_width) + 1
|
||||
range_end = min(int((bucket + 1) * bucket_width) + 1, len(points) - 1)
|
||||
anchor = points[anchor_index]
|
||||
selected_index = range_start
|
||||
maximum_area = -1.0
|
||||
for index in range(range_start, max(range_start + 1, range_end)):
|
||||
point = points[index]
|
||||
area = abs(
|
||||
(anchor["step"] - average_x) * (point["value"] - anchor["value"])
|
||||
- (anchor["step"] - point["step"]) * (average_y - anchor["value"])
|
||||
)
|
||||
if area > maximum_area:
|
||||
maximum_area = area
|
||||
selected_index = index
|
||||
sampled.append(points[selected_index])
|
||||
anchor_index = selected_index
|
||||
sampled.append(points[-1])
|
||||
return sampled
|
||||
|
||||
|
||||
class TuningStorage:
|
||||
def __init__(self, path: Path):
|
||||
self.path = path.expanduser().resolve()
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self.local = threading.local()
|
||||
self._migrate()
|
||||
|
||||
def connection(self) -> sqlite3.Connection:
|
||||
connection = getattr(self.local, "connection", None)
|
||||
if connection is None:
|
||||
connection = sqlite3.connect(self.path, timeout=10, isolation_level=None)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys=ON")
|
||||
connection.execute("PRAGMA journal_mode=WAL")
|
||||
connection.execute("PRAGMA busy_timeout=10000")
|
||||
self.local.connection = connection
|
||||
return connection
|
||||
|
||||
@contextmanager
|
||||
def transaction(self) -> Iterator[sqlite3.Connection]:
|
||||
connection = self.connection()
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
yield connection
|
||||
connection.execute("COMMIT")
|
||||
except Exception:
|
||||
connection.execute("ROLLBACK")
|
||||
raise
|
||||
|
||||
def _migrate(self) -> None:
|
||||
connection = self.connection()
|
||||
connection.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations(version INTEGER PRIMARY KEY);
|
||||
CREATE TABLE IF NOT EXISTS sessions(
|
||||
id TEXT PRIMARY KEY, state TEXT NOT NULL, mode TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL, updated_at TEXT NOT NULL,
|
||||
config_json TEXT NOT NULL, objective_json TEXT NOT NULL,
|
||||
message TEXT NOT NULL, current_trial_id TEXT, best_trial_id TEXT,
|
||||
consecutive_no_improve INTEGER NOT NULL DEFAULT 0,
|
||||
fallback_enabled INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS trials(
|
||||
id TEXT PRIMARY KEY,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
number INTEGER NOT NULL, state TEXT NOT NULL, rung INTEGER NOT NULL,
|
||||
target_iterations INTEGER NOT NULL, reward_config_json TEXT NOT NULL,
|
||||
proposal_id TEXT, run_dir TEXT NOT NULL, checkpoint_path TEXT,
|
||||
policy_path TEXT, evaluation_json TEXT, score REAL, eligible INTEGER,
|
||||
created_at TEXT NOT NULL, started_at TEXT, ended_at TEXT, message TEXT NOT NULL,
|
||||
UNIQUE(session_id, number, rung)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS proposals(
|
||||
id TEXT PRIMARY KEY,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
base_trial_id TEXT, state TEXT NOT NULL, source TEXT NOT NULL,
|
||||
patch_json TEXT NOT NULL, rationale TEXT NOT NULL,
|
||||
expected_json TEXT, confidence REAL NOT NULL,
|
||||
created_at TEXT NOT NULL, decided_at TEXT, feedback TEXT
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS metric_points(
|
||||
trial_id TEXT NOT NULL REFERENCES trials(id) ON DELETE CASCADE,
|
||||
tag TEXT NOT NULL, step INTEGER NOT NULL,
|
||||
wall_time REAL NOT NULL, value REAL NOT NULL,
|
||||
PRIMARY KEY(trial_id, tag, step)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS audit_events(
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
event_type TEXT NOT NULL, payload_json TEXT NOT NULL, created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS presets(
|
||||
id TEXT PRIMARY KEY, name TEXT NOT NULL UNIQUE, session_id TEXT NOT NULL,
|
||||
trial_id TEXT NOT NULL, reward_config_json TEXT NOT NULL, created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_trials_session ON trials(session_id, number, rung);
|
||||
CREATE INDEX IF NOT EXISTS idx_proposals_session ON proposals(session_id, created_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_metrics_trial_tag ON metric_points(trial_id, tag, step);
|
||||
"""
|
||||
)
|
||||
connection.execute(
|
||||
"INSERT OR IGNORE INTO schema_migrations(version) VALUES (?)", (SCHEMA_VERSION,)
|
||||
)
|
||||
|
||||
def recover_interrupted(self) -> None:
|
||||
at = now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"UPDATE trials SET state='interrupted', ended_at=?, "
|
||||
"message='服务重启中断,等待显式恢复' "
|
||||
"WHERE state IN ('training','evaluating')",
|
||||
(at,),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET state='interrupted', updated_at=?, "
|
||||
"message='服务重启中断,可从完整 checkpoint 恢复' "
|
||||
"WHERE state IN ('running','evaluating')",
|
||||
(at,),
|
||||
)
|
||||
|
||||
def create_session(self, mode: str, config: dict, objective: dict, fallback: bool) -> dict:
|
||||
session_id, at = uuid.uuid4().hex, now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"INSERT INTO sessions("
|
||||
"id,state,mode,created_at,updated_at,config_json,objective_json,"
|
||||
"message,fallback_enabled) "
|
||||
"VALUES (?, 'queued', ?, ?, ?, ?, ?, '等待基线训练', ?)",
|
||||
(session_id, mode, at, at, _json(config), _json(objective), int(fallback)),
|
||||
)
|
||||
connection.execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
(session_id, "session_created", _json({"mode": mode}), at),
|
||||
)
|
||||
return self.get_session(session_id)
|
||||
|
||||
def _session(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"state": row["state"],
|
||||
"mode": row["mode"],
|
||||
"createdAt": row["created_at"],
|
||||
"updatedAt": row["updated_at"],
|
||||
"config": _decode(row["config_json"]),
|
||||
"objectiveWeights": _decode(row["objective_json"]),
|
||||
"message": row["message"],
|
||||
"currentTrialId": row["current_trial_id"],
|
||||
"bestTrialId": row["best_trial_id"],
|
||||
"consecutiveNoImprove": row["consecutive_no_improve"],
|
||||
"fallbackEnabled": bool(row["fallback_enabled"]),
|
||||
}
|
||||
|
||||
def get_session(self, session_id: str) -> dict:
|
||||
row = (
|
||||
self.connection().execute("SELECT * FROM sessions WHERE id=?", (session_id,)).fetchone()
|
||||
)
|
||||
if row is None:
|
||||
raise KeyError(session_id)
|
||||
return self._session(row)
|
||||
|
||||
def list_sessions(self, limit: int = 50) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM sessions ORDER BY created_at DESC LIMIT ?", (limit,))
|
||||
.fetchall()
|
||||
)
|
||||
return [self._session(row) for row in rows]
|
||||
|
||||
def update_session(self, session_id: str, **changes: Any) -> bool:
|
||||
columns = {
|
||||
"state": "state",
|
||||
"message": "message",
|
||||
"current_trial_id": "current_trial_id",
|
||||
"best_trial_id": "best_trial_id",
|
||||
"consecutive_no_improve": "consecutive_no_improve",
|
||||
}
|
||||
values, assignments = [], []
|
||||
for key, value in changes.items():
|
||||
if key not in columns:
|
||||
raise ValueError(key)
|
||||
assignments.append(f"{columns[key]}=?")
|
||||
values.append(value)
|
||||
assignments.append("updated_at=?")
|
||||
values.extend((now_iso(), session_id))
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE sessions SET {', '.join(assignments)} WHERE id=?", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def create_trial(
|
||||
self,
|
||||
session_id: str,
|
||||
number: int,
|
||||
rung: int,
|
||||
target: int,
|
||||
reward_config: dict,
|
||||
proposal_id: str | None,
|
||||
run_dir: str,
|
||||
) -> dict:
|
||||
trial_id, at = uuid.uuid4().hex, now_iso()
|
||||
with self.transaction() as connection:
|
||||
connection.execute(
|
||||
"INSERT INTO trials("
|
||||
"id,session_id,number,state,rung,target_iterations,reward_config_json,"
|
||||
"proposal_id,run_dir,created_at,message) "
|
||||
"VALUES (?,?,?,'queued',?,?,?,?,?,?,'等待训练')",
|
||||
(
|
||||
trial_id,
|
||||
session_id,
|
||||
number,
|
||||
rung,
|
||||
target,
|
||||
_json(reward_config),
|
||||
proposal_id,
|
||||
run_dir,
|
||||
at,
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET current_trial_id=?,updated_at=? WHERE id=?",
|
||||
(trial_id, at, session_id),
|
||||
)
|
||||
return self.get_trial(trial_id)
|
||||
|
||||
def _trial(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"sessionId": row["session_id"],
|
||||
"number": row["number"],
|
||||
"state": row["state"],
|
||||
"rung": row["rung"],
|
||||
"targetIterations": row["target_iterations"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"proposalId": row["proposal_id"],
|
||||
"runDir": row["run_dir"],
|
||||
"checkpointPath": row["checkpoint_path"],
|
||||
"policyPath": row["policy_path"],
|
||||
"evaluation": _decode(row["evaluation_json"]),
|
||||
"score": row["score"],
|
||||
"eligible": None if row["eligible"] is None else bool(row["eligible"]),
|
||||
"createdAt": row["created_at"],
|
||||
"startedAt": row["started_at"],
|
||||
"endedAt": row["ended_at"],
|
||||
"message": row["message"],
|
||||
}
|
||||
|
||||
def get_trial(self, trial_id: str) -> dict:
|
||||
row = self.connection().execute("SELECT * FROM trials WHERE id=?", (trial_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(trial_id)
|
||||
return self._trial(row)
|
||||
|
||||
def list_trials(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM trials WHERE session_id=? ORDER BY number,rung", (session_id,))
|
||||
.fetchall()
|
||||
)
|
||||
return [self._trial(row) for row in rows]
|
||||
|
||||
def delete_trial(self, trial_id: str) -> bool:
|
||||
cursor = self.connection().execute(
|
||||
"DELETE FROM trials WHERE id=? AND state='interrupted'", (trial_id,)
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def update_trial(self, trial_id: str, **changes: Any) -> bool:
|
||||
columns = {
|
||||
"state": "state",
|
||||
"message": "message",
|
||||
"checkpoint_path": "checkpoint_path",
|
||||
"policy_path": "policy_path",
|
||||
"score": "score",
|
||||
"eligible": "eligible",
|
||||
"started_at": "started_at",
|
||||
"ended_at": "ended_at",
|
||||
"evaluation": "evaluation_json",
|
||||
}
|
||||
values, assignments = [], []
|
||||
for key, value in changes.items():
|
||||
if key not in columns:
|
||||
raise ValueError(key)
|
||||
if key == "evaluation":
|
||||
value = _json(value)
|
||||
if key == "eligible":
|
||||
value = int(value)
|
||||
assignments.append(f"{columns[key]}=?")
|
||||
values.append(value)
|
||||
values.append(trial_id)
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE trials SET {', '.join(assignments)} WHERE id=?", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def create_proposal(
|
||||
self,
|
||||
session_id: str,
|
||||
base_trial_id: str | None,
|
||||
patch: dict,
|
||||
rationale: str,
|
||||
expected: Any,
|
||||
confidence: float,
|
||||
source: str = "agent",
|
||||
) -> dict:
|
||||
proposal_id, at = uuid.uuid4().hex, now_iso()
|
||||
self.connection().execute(
|
||||
"INSERT INTO proposals("
|
||||
"id,session_id,base_trial_id,state,source,patch_json,rationale,"
|
||||
"expected_json,confidence,created_at) "
|
||||
"VALUES (?,?,?,'pending',?,?,?,?,?,?)",
|
||||
(
|
||||
proposal_id,
|
||||
session_id,
|
||||
base_trial_id,
|
||||
source,
|
||||
_json(patch),
|
||||
rationale,
|
||||
_json(expected),
|
||||
confidence,
|
||||
at,
|
||||
),
|
||||
)
|
||||
return self.get_proposal(proposal_id)
|
||||
|
||||
def _proposal(self, row: sqlite3.Row) -> dict:
|
||||
return {
|
||||
"id": row["id"],
|
||||
"sessionId": row["session_id"],
|
||||
"baseTrialId": row["base_trial_id"],
|
||||
"state": row["state"],
|
||||
"source": row["source"],
|
||||
"patch": _decode(row["patch_json"]),
|
||||
"rationale": row["rationale"],
|
||||
"expectedImpact": _decode(row["expected_json"]),
|
||||
"confidence": row["confidence"],
|
||||
"createdAt": row["created_at"],
|
||||
"decidedAt": row["decided_at"],
|
||||
"feedback": row["feedback"],
|
||||
}
|
||||
|
||||
def get_proposal(self, proposal_id: str) -> dict:
|
||||
row = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM proposals WHERE id=?", (proposal_id,))
|
||||
.fetchone()
|
||||
)
|
||||
if row is None:
|
||||
raise KeyError(proposal_id)
|
||||
return self._proposal(row)
|
||||
|
||||
def list_proposals(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
"SELECT * FROM proposals WHERE session_id=? ORDER BY created_at", (session_id,)
|
||||
)
|
||||
.fetchall()
|
||||
)
|
||||
return [self._proposal(row) for row in rows]
|
||||
|
||||
def decide_proposal(
|
||||
self, proposal_id: str, state: str, feedback: str | None, patch: dict | None = None
|
||||
) -> bool:
|
||||
at = now_iso()
|
||||
assignments, values = ["state=?", "feedback=?", "decided_at=?"], [state, feedback, at]
|
||||
if patch is not None:
|
||||
assignments.append("patch_json=?")
|
||||
values.append(_json(patch))
|
||||
values.extend((proposal_id,))
|
||||
cursor = self.connection().execute(
|
||||
f"UPDATE proposals SET {', '.join(assignments)} WHERE id=? AND state='pending'", values
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def insert_metrics(self, trial_id: str, points: list[tuple[str, int, float, float]]) -> None:
|
||||
self.connection().executemany(
|
||||
"INSERT INTO metric_points(trial_id,tag,step,wall_time,value) "
|
||||
"VALUES (?,?,?,?,?) ON CONFLICT(trial_id,tag,step) DO UPDATE SET "
|
||||
"wall_time=excluded.wall_time,value=excluded.value",
|
||||
[(trial_id, *point) for point in points],
|
||||
)
|
||||
|
||||
def metrics(
|
||||
self, trial_id: str, tags: list[str] | None = None, max_points: int = 1000
|
||||
) -> list[dict]:
|
||||
parameters: list[Any] = [trial_id]
|
||||
clause = "trial_id=?"
|
||||
if tags:
|
||||
clause += f" AND tag IN ({','.join('?' for _ in tags)})"
|
||||
parameters.extend(tags)
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
f"SELECT tag,step,wall_time,value FROM metric_points "
|
||||
f"WHERE {clause} ORDER BY tag,step",
|
||||
parameters,
|
||||
)
|
||||
.fetchall()
|
||||
)
|
||||
grouped: dict[str, list[dict]] = {}
|
||||
for row in rows:
|
||||
grouped.setdefault(row["tag"], []).append(
|
||||
{"step": row["step"], "wallTime": row["wall_time"], "value": row["value"]}
|
||||
)
|
||||
series = []
|
||||
for tag, values in grouped.items():
|
||||
if len(values) > max_points:
|
||||
values = _lttb(values, max_points)
|
||||
series.append({"tag": tag, "points": values})
|
||||
return series
|
||||
|
||||
def audit(self, session_id: str, event_type: str, payload: Any) -> None:
|
||||
self.connection().execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
(session_id, event_type, _json(payload), now_iso()),
|
||||
)
|
||||
|
||||
def audit_events(self, session_id: str) -> list[dict]:
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM audit_events WHERE session_id=? ORDER BY id", (session_id,))
|
||||
.fetchall()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": row["id"],
|
||||
"type": row["event_type"],
|
||||
"payload": _decode(row["payload_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def save_preset(self, name: str, session_id: str, trial_id: str, reward_config: dict) -> dict:
|
||||
preset_id, at = uuid.uuid4().hex, now_iso()
|
||||
self.connection().execute(
|
||||
"INSERT INTO presets(id,name,session_id,trial_id,reward_config_json,created_at) "
|
||||
"VALUES (?,?,?,?,?,?)",
|
||||
(preset_id, name, session_id, trial_id, _json(reward_config), at),
|
||||
)
|
||||
return {
|
||||
"id": preset_id,
|
||||
"name": name,
|
||||
"sessionId": session_id,
|
||||
"trialId": trial_id,
|
||||
"rewardConfig": reward_config,
|
||||
"createdAt": at,
|
||||
}
|
||||
|
||||
def get_preset(self, preset_id: str) -> dict:
|
||||
row = self.connection().execute("SELECT * FROM presets WHERE id=?", (preset_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(preset_id)
|
||||
return {
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
|
||||
def list_presets(self) -> list[dict]:
|
||||
rows = (
|
||||
self.connection().execute("SELECT * FROM presets ORDER BY created_at DESC").fetchall()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Lazy Optuna study integration used for durable trial history and pruning metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
class OptunaStudies:
|
||||
def __init__(self, root: Path):
|
||||
self.path = (root / "optuna.sqlite3").resolve()
|
||||
self.url = f"sqlite:///{self.path}"
|
||||
|
||||
def _study(self, session_id: str):
|
||||
import optuna
|
||||
|
||||
optuna.logging.set_verbosity(optuna.logging.WARNING)
|
||||
return optuna.create_study(
|
||||
study_name=f"reward-tuning-{session_id}",
|
||||
storage=self.url,
|
||||
direction="maximize",
|
||||
load_if_exists=True,
|
||||
pruner=optuna.pruners.SuccessiveHalvingPruner(
|
||||
min_resource=300, reduction_factor=3, min_early_stopping_rate=0
|
||||
),
|
||||
)
|
||||
|
||||
def record(
|
||||
self,
|
||||
session_id: str,
|
||||
reward_config: dict[str, Any],
|
||||
score: float,
|
||||
eligible: bool,
|
||||
rung: int,
|
||||
) -> int:
|
||||
"""Record an externally proposed Agent config through Optuna ask/tell."""
|
||||
import optuna
|
||||
|
||||
study = self._study(session_id)
|
||||
trial = study.ask()
|
||||
trial.set_user_attr("reward_config", reward_config)
|
||||
trial.set_user_attr("eligible", eligible)
|
||||
trial.set_user_attr("rung", rung)
|
||||
state = optuna.trial.TrialState.COMPLETE if eligible else optuna.trial.TrialState.PRUNED
|
||||
study.tell(trial, score if eligible else None, state=state)
|
||||
return trial.number
|
||||
|
||||
def summary(self, session_id: str) -> dict[str, Any]:
|
||||
study = self._study(session_id)
|
||||
completed = [trial for trial in study.trials if trial.value is not None]
|
||||
return {
|
||||
"studyName": study.study_name,
|
||||
"trialCount": len(study.trials),
|
||||
"bestValue": max((trial.value for trial in completed), default=None),
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
"""TensorBoard scalar ingestion with optional dependency isolation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from .storage import TuningStorage
|
||||
|
||||
|
||||
class TensorboardUnavailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def ingest_scalars(storage: TuningStorage, trial_id: str, log_dir: Path) -> int:
|
||||
"""Reload all scalar events and idempotently upsert them into SQLite."""
|
||||
try:
|
||||
from tensorboard.backend.event_processing.event_accumulator import EventAccumulator
|
||||
except ImportError as error:
|
||||
raise TensorboardUnavailable(
|
||||
"缺少 tensorboard,请安装 training_server/requirements.txt"
|
||||
) from error
|
||||
if not log_dir.is_dir():
|
||||
return 0
|
||||
accumulator = EventAccumulator(str(log_dir), size_guidance={"scalars": 0})
|
||||
try:
|
||||
accumulator.Reload()
|
||||
except (OSError, ValueError):
|
||||
return 0
|
||||
points: list[tuple[str, int, float, float]] = []
|
||||
for tag in accumulator.Tags().get("scalars", []):
|
||||
for event in accumulator.Scalars(tag):
|
||||
points.append((tag, int(event.step), float(event.wall_time), float(event.value)))
|
||||
if points:
|
||||
storage.insert_metrics(trial_id, points)
|
||||
return len(points)
|
||||
Reference in New Issue
Block a user