feat(tuning): release V0.8.2 Agent 界面重构
This commit is contained in:
@@ -57,7 +57,7 @@ python training_server/server.py \
|
||||
|
||||
在主工作台连接训练服务后,点击“打开自调参 Agent 工作台”。默认预算为 12 个唯一配置:所有配置先训练 300 iterations,前 4 名续训到 900,前 2 名续训到 2000;默认使用 GPU 0 和 4096 个并行环境。首次使用建议先降低为 256–512 environments 做 smoke test。
|
||||
|
||||
固定评估使用站立、前进/侧移、转向和组合命令以及 3 个固定 seed。最终分数不直接使用可被权重放大的总 reward,而由速度跟踪 35%、动作平滑 20%、姿态稳定 15%、减少跌倒 15%、足端滑移 10%、能耗 5% 的权重无关指标组成。跌倒率高于基线 2% 或速度误差恶化超过 5% 的 trial 不晋级。逐轮审批模式会自动运行基线,之后每条 Agent 建议都等待批准、修改后批准或拒绝反馈。
|
||||
固定评估使用站立、前进/侧移、转向和组合命令以及 3 个固定 seed。最终分数不直接使用可被权重放大的总 reward,而由速度跟踪 35%、动作平滑 20%、姿态稳定 15%、减少跌倒 15%、足端滑移 10%、能耗 5% 的权重无关指标组成。跌倒率高于基线 2% 或速度误差恶化超过 5% 的 trial 不晋级。逐轮审批模式会自动运行基线,之后每条 Agent 建议都等待批准、修改后批准或拒绝反馈;运行中的 session 也可在逐轮审批和全自动之间切换,切到全自动时会批准当前待处理建议。
|
||||
|
||||
最佳结果保存为不可变 preset,可在普通训练面板的“奖励配置”中选择,也可导出 JSON;不会覆盖仓库里的 Python 默认奖励配置。
|
||||
|
||||
@@ -71,12 +71,16 @@ python training_server/server.py \
|
||||
- `GET /api/tuning/capabilities`、`POST /api/tuning/agent/test`:检查/测试 Agent;
|
||||
- `GET|POST /api/tuning/sessions`、`GET|DELETE /api/tuning/sessions/{id}`:列出、创建、查询、停止 session;
|
||||
- `POST /api/tuning/sessions/{id}/pause|resume`:暂停后续调度或恢复;
|
||||
- `POST /api/tuning/sessions/{id}/mode`:运行时切换 `automatic`/`approval` 模式;
|
||||
- `PUT /api/tuning/sessions/{id}/constraints`:以 revision CAS 保存参数固定值/工程上下限,服务端在 Agent、fallback 与人工修改三条路径统一强制;
|
||||
- `POST /api/tuning/sessions/{id}/step`:发放且只消费一个 Trial 调度令牌,完成训练与固定评估后重新暂停;
|
||||
- `POST /api/tuning/sessions/{id}/rollback`:把同 Session 内已完成且通过安全门槛的 Trial 设为非破坏性后续基准,可同时验证其 checkpoint;
|
||||
- `POST /api/tuning/sessions/{id}/proposals/{proposalId}/approve|reject`:审批、修改或拒绝建议;
|
||||
- `GET /api/tuning/sessions/{id}/trials/{trialId}/metrics`:查询降采样 scalar;
|
||||
- `GET /api/tuning/sessions/{id}/trials/{trialId}/metrics?afterStep=N`:查询降采样或增量 scalar;
|
||||
- `GET /api/tuning/sessions/{id}/artifacts/best/policy.onnx`:下载最佳策略;
|
||||
- `GET /api/tuning/presets`:列出可供普通训练复用的最佳奖励 preset。
|
||||
|
||||
普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态及产物持久化在 `logs/auto_tuning/`。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、每轮最多 4 项及 `0.5×–2×` 变化率校验。
|
||||
普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态、调度令牌、参数护栏、回滚基准及产物持久化在 `logs/auto_tuning/`。候选配置数可在 1–100 间设置(包含基线配置,仍受连续无提升早停约束)。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、Session 护栏、每轮最多 4 项及 `0.5×–2×` 变化率校验。
|
||||
|
||||
## 测试
|
||||
|
||||
|
||||
@@ -27,6 +27,8 @@ from urllib.parse import parse_qs, unquote, urlsplit
|
||||
|
||||
from tuning.manager import TuningError, TuningManager
|
||||
from tuning.process import GpuLease, ResourceBusyError
|
||||
from tuning.schema import RewardConfigError
|
||||
from tuning.scoring import EvaluationError
|
||||
|
||||
VERSION = "0.4.0"
|
||||
# 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。
|
||||
@@ -498,7 +500,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
self._json(HTTPStatus.NOT_FOUND, {"error": "调参 session、trial 或 proposal 不存在"})
|
||||
elif isinstance(error, ResourceBusyError):
|
||||
self._json(HTTPStatus.CONFLICT, {"error": str(error)})
|
||||
elif isinstance(error, TuningError):
|
||||
elif isinstance(error, (TuningError, RewardConfigError, EvaluationError)):
|
||||
self._json(HTTPStatus.BAD_REQUEST, {"error": str(error)})
|
||||
else:
|
||||
self._json(
|
||||
@@ -552,7 +554,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
self._ensure_origin()
|
||||
self.send_response(HTTPStatus.NO_CONTENT)
|
||||
self._cors()
|
||||
self.send_header("Access-Control-Allow-Methods", "GET, POST, DELETE, OPTIONS")
|
||||
self.send_header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
self.send_header("Access-Control-Allow-Headers", "Authorization, Content-Type")
|
||||
self.send_header("Access-Control-Max-Age", "600")
|
||||
self.end_headers()
|
||||
@@ -591,14 +593,18 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
tags = [tag for value in query.get("tags", []) for tag in value.split(",") if tag]
|
||||
try:
|
||||
max_points = int(query.get("maxPoints", ["1000"])[0])
|
||||
after_raw = query.get("afterStep", [None])[0]
|
||||
after_step = int(after_raw) if after_raw is not None else None
|
||||
except ValueError as error:
|
||||
raise TuningError("maxPoints 必须是整数") from error
|
||||
raise TuningError("maxPoints/afterStep 必须是整数") from error
|
||||
if not 10 <= max_points <= 5000:
|
||||
raise TuningError("maxPoints 必须在 10–5000 之间")
|
||||
if after_step is not None and after_step < -1:
|
||||
raise TuningError("afterStep 不能小于 -1")
|
||||
self._json(
|
||||
HTTPStatus.OK,
|
||||
self.tuning_manager.metrics(
|
||||
match.group(1), match.group(2), tags or None, max_points
|
||||
match.group(1), match.group(2), tags or None, max_points, after_step
|
||||
),
|
||||
)
|
||||
return
|
||||
@@ -640,6 +646,22 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
)
|
||||
self._json(HTTPStatus.ACCEPTED, action(match.group(1)))
|
||||
return
|
||||
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})/(step|rollback)", path)
|
||||
if match:
|
||||
action = (
|
||||
self.tuning_manager.step
|
||||
if match.group(2) == "step"
|
||||
else self.tuning_manager.rollback
|
||||
)
|
||||
self._json(HTTPStatus.ACCEPTED, action(match.group(1), self._payload()))
|
||||
return
|
||||
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})/mode", path)
|
||||
if match:
|
||||
self._json(
|
||||
HTTPStatus.ACCEPTED,
|
||||
self.tuning_manager.set_mode(match.group(1), self._payload()),
|
||||
)
|
||||
return
|
||||
match = re.fullmatch(
|
||||
r"/api/tuning/sessions/([0-9a-f]{32})/proposals/([0-9a-f]{32})/(approve|reject)",
|
||||
path,
|
||||
@@ -658,6 +680,21 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
except Exception as error:
|
||||
self._error(error)
|
||||
|
||||
def do_PUT(self) -> None:
|
||||
try:
|
||||
self._ensure_request()
|
||||
path = urlsplit(self.path).path
|
||||
match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})/constraints", path)
|
||||
if match:
|
||||
self._json(
|
||||
HTTPStatus.ACCEPTED,
|
||||
self.tuning_manager.set_constraints(match.group(1), self._payload()),
|
||||
)
|
||||
return
|
||||
raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在")
|
||||
except Exception as error:
|
||||
self._error(error)
|
||||
|
||||
def do_DELETE(self) -> None:
|
||||
try:
|
||||
self._ensure_request()
|
||||
|
||||
@@ -12,6 +12,7 @@ from tuning.schema import ( # noqa: E402
|
||||
RewardConfigError,
|
||||
merge_proposal,
|
||||
validate_configuration,
|
||||
validate_constraints,
|
||||
validate_proposal,
|
||||
)
|
||||
from tuning.scoring import ( # noqa: E402
|
||||
@@ -19,7 +20,7 @@ from tuning.scoring import ( # noqa: E402
|
||||
EvaluationError,
|
||||
score_evaluation,
|
||||
)
|
||||
from tuning.storage import TuningStorage # noqa: E402
|
||||
from tuning.storage import StorageConflict, TuningStorage # noqa: E402
|
||||
|
||||
|
||||
class RewardSchemaTest(unittest.TestCase):
|
||||
@@ -69,6 +70,33 @@ class RewardSchemaTest(unittest.TestCase):
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
)
|
||||
|
||||
def test_session_constraints_reject_unknown_out_of_range_and_fixed_changes(self):
|
||||
constraints = validate_constraints(
|
||||
{
|
||||
"weights.track_linear_velocity": {"kind": "fixed", "value": 1.0},
|
||||
"params.foot_gait.period": {"kind": "range", "min": 0.5, "max": 0.7},
|
||||
}
|
||||
)
|
||||
validate_proposal(
|
||||
{"params": {"foot_gait.period": 0.65}},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
constraints,
|
||||
)
|
||||
with self.assertRaisesRegex(RewardConfigError, "已固定"):
|
||||
validate_proposal(
|
||||
{"weights": {"track_linear_velocity": 1.1}},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
constraints,
|
||||
)
|
||||
with self.assertRaisesRegex(RewardConfigError, "工程锁定范围"):
|
||||
validate_proposal(
|
||||
{"params": {"foot_gait.period": 0.75}},
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
constraints,
|
||||
)
|
||||
with self.assertRaisesRegex(RewardConfigError, "未知参数约束"):
|
||||
validate_constraints({"weights.not_allowed": {"kind": "fixed", "value": 1.0}})
|
||||
|
||||
|
||||
class ScoringTest(unittest.TestCase):
|
||||
baseline = {
|
||||
@@ -174,8 +202,24 @@ class StorageTest(unittest.TestCase):
|
||||
self.assertEqual(sampled[0]["step"], 0)
|
||||
self.assertEqual(sampled[-1]["step"], 99)
|
||||
self.assertIn(50.0, [point["value"] for point in sampled])
|
||||
control = self.storage.replace_constraints(
|
||||
session["id"],
|
||||
0,
|
||||
{"weights.pose": {"kind": "range", "min": 0.5, "max": 1.5}},
|
||||
)
|
||||
self.assertEqual(control["constraintsRevision"], 1)
|
||||
with self.assertRaises(StorageConflict):
|
||||
self.storage.replace_constraints(session["id"], 0, {})
|
||||
self.storage.grant_dispatch_token(session["id"])
|
||||
with self.assertRaises(StorageConflict):
|
||||
self.storage.grant_dispatch_token(session["id"])
|
||||
self.assertTrue(self.storage.use_dispatch_token(session["id"]))
|
||||
self.assertFalse(self.storage.use_dispatch_token(session["id"]))
|
||||
incremental = self.storage.metrics(trial["id"], max_points=100, after_step=90)[0]
|
||||
self.assertEqual(incremental["points"][0]["step"], 91)
|
||||
reopened = TuningStorage(self.storage.path)
|
||||
self.assertEqual(reopened.get_session(session["id"])["mode"], "approval")
|
||||
self.assertEqual(reopened.get_control(session["id"])["constraintsRevision"], 1)
|
||||
|
||||
def test_recovery_marks_inflight_records(self):
|
||||
session = self.storage.create_session(
|
||||
|
||||
@@ -148,6 +148,8 @@ class TuningManagerTest(unittest.TestCase):
|
||||
self.fail("session did not wait for approval")
|
||||
proposal = detail["proposals"][-1]
|
||||
patch = {"weights": {"pose": 1.1}, "params": {}}
|
||||
with self.assertRaisesRegex(Exception, "feedback"):
|
||||
self.manager.approve(session["id"], proposal["id"], {"feedback": {}})
|
||||
approved = self.manager.approve(
|
||||
session["id"], proposal["id"], {"feedback": "ok", "patch": patch}
|
||||
)
|
||||
@@ -155,6 +157,34 @@ class TuningManagerTest(unittest.TestCase):
|
||||
self.manager.cancel(session["id"])
|
||||
self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled")
|
||||
|
||||
def test_runtime_mode_switch_auto_approves_pending_proposal(self):
|
||||
session = self.manager.create(self.payload("approval"))
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if detail["state"] == "awaiting_approval":
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("session did not wait for approval")
|
||||
changed = self.manager.set_mode(session["id"], {"mode": "automatic"})
|
||||
self.assertEqual(changed["mode"], "automatic")
|
||||
self.assertEqual(changed["proposals"][-1]["state"], "approved")
|
||||
completed = self.wait_terminal(session["id"])
|
||||
self.assertEqual(completed["state"], "succeeded", completed["message"])
|
||||
|
||||
def test_trial_count_is_user_configurable(self):
|
||||
payload = self.payload()
|
||||
payload["trialCount"] = 1
|
||||
_, config, _, _ = self.manager.parse_create(payload)
|
||||
self.assertEqual(config["trialCount"], 1)
|
||||
payload["trialCount"] = 100
|
||||
_, config, _, _ = self.manager.parse_create(payload)
|
||||
self.assertEqual(config["trialCount"], 100)
|
||||
payload["trialCount"] = 101
|
||||
with self.assertRaisesRegex(Exception, "trialCount"):
|
||||
self.manager.parse_create(payload)
|
||||
|
||||
def test_resume_discards_only_interrupted_trial_and_continues(self):
|
||||
mode, config, objective, fallback = self.manager.parse_create(self.payload())
|
||||
session = self.manager.storage.create_session(mode, config, objective, fallback)
|
||||
@@ -178,8 +208,10 @@ class TuningManagerTest(unittest.TestCase):
|
||||
"trial-001-rung-0",
|
||||
)
|
||||
self.manager.storage.update_trial(interrupted["id"], state="interrupted")
|
||||
self.manager.storage.reset_step_gate(session["id"])
|
||||
self.manager.storage.update_session(session["id"], state="interrupted")
|
||||
self.manager.resume(session["id"])
|
||||
self.assertEqual(self.manager.storage.get_control(session["id"])["runPolicy"], "continuous")
|
||||
completed = self.wait_terminal(session["id"])
|
||||
self.assertEqual(completed["state"], "succeeded", completed["message"])
|
||||
self.assertNotIn(interrupted["id"], [trial["id"] for trial in completed["trials"]])
|
||||
@@ -190,6 +222,129 @@ class TuningManagerTest(unittest.TestCase):
|
||||
self.manager.parse_create({"taskId": "Other"})
|
||||
self.assertEqual(self.manager.test_agent()["model"], "fake")
|
||||
|
||||
def test_pause_closes_persistent_dispatch_gate_at_trial_boundary(self):
|
||||
mode, config, objective, fallback = self.manager.parse_create(self.payload())
|
||||
session = self.manager.storage.create_session(mode, config, objective, fallback)
|
||||
self.manager.storage.update_session(session["id"], state="running")
|
||||
self.manager.storage.grant_dispatch_token(session["id"])
|
||||
|
||||
paused = self.manager.pause(session["id"])
|
||||
|
||||
self.assertEqual(paused["state"], "paused")
|
||||
self.assertEqual(paused["control"]["runPolicy"], "step")
|
||||
self.assertEqual(paused["control"]["dispatchTokens"], 0)
|
||||
|
||||
def test_step_token_executes_exactly_one_trial_then_pauses(self):
|
||||
session = self.manager.create(self.payload("approval"))
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if detail["state"] == "awaiting_approval":
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("session did not wait for approval")
|
||||
baseline_count = len([trial for trial in detail["trials"] if trial["state"] == "completed"])
|
||||
stepped = self.manager.step(session["id"], {"count": 1})
|
||||
self.assertEqual(stepped["control"]["runPolicy"], "step")
|
||||
self.assertEqual(stepped["control"]["dispatchTokens"], 1)
|
||||
proposal = stepped["proposals"][-1]
|
||||
self.manager.approve(session["id"], proposal["id"], {})
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
completed_count = len(
|
||||
[trial for trial in detail["trials"] if trial["state"] == "completed"]
|
||||
)
|
||||
if detail["state"] == "paused" and completed_count == baseline_count + 1:
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("single-step trial did not pause at the next boundary")
|
||||
time.sleep(0.05)
|
||||
self.assertEqual(
|
||||
len(
|
||||
[
|
||||
trial
|
||||
for trial in self.manager.detail(session["id"])["trials"]
|
||||
if trial["state"] == "completed"
|
||||
]
|
||||
),
|
||||
baseline_count + 1,
|
||||
)
|
||||
self.manager.cancel(session["id"])
|
||||
self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled")
|
||||
|
||||
def test_constraints_are_revisioned_and_enforced_during_approval(self):
|
||||
session = self.manager.create(self.payload("approval"))
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if detail["state"] == "awaiting_approval":
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("session did not wait for approval")
|
||||
proposal = detail["proposals"][-1]
|
||||
constrained = self.manager.set_constraints(
|
||||
session["id"],
|
||||
{
|
||||
"revision": 0,
|
||||
"constraints": {"weights.track_linear_velocity": {"kind": "fixed", "value": 1.0}},
|
||||
},
|
||||
)
|
||||
self.assertEqual(constrained["control"]["constraintsRevision"], 1)
|
||||
with self.assertRaisesRegex(Exception, "已固定"):
|
||||
self.manager.approve(
|
||||
session["id"],
|
||||
proposal["id"],
|
||||
{"patch": {"weights": {"track_linear_velocity": 1.1}, "params": {}}},
|
||||
)
|
||||
with self.assertRaisesRegex(Exception, "revision"):
|
||||
self.manager.set_constraints(session["id"], {"revision": 0, "constraints": {}})
|
||||
self.manager.approve(session["id"], proposal["id"], {})
|
||||
self.manager.cancel(session["id"])
|
||||
self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled")
|
||||
|
||||
def test_rollback_uses_safe_completed_trial_as_next_proposal_base(self):
|
||||
session = self.manager.create(self.payload("approval"))
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if detail["state"] == "awaiting_approval":
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("session did not wait for approval")
|
||||
baseline = detail["trials"][0]
|
||||
old_proposal = detail["proposals"][-1]
|
||||
rolled_back = self.manager.rollback(
|
||||
session["id"], {"trialId": baseline["id"], "checkpoint": True}
|
||||
)
|
||||
self.assertEqual(rolled_back["state"], "paused")
|
||||
self.assertEqual(rolled_back["control"]["activeBaseTrialId"], baseline["id"])
|
||||
self.assertEqual(
|
||||
next(item for item in rolled_back["proposals"] if item["id"] == old_proposal["id"])[
|
||||
"state"
|
||||
],
|
||||
"rejected",
|
||||
)
|
||||
self.manager.step(session["id"], {"count": 1})
|
||||
deadline = time.monotonic() + 3
|
||||
while time.monotonic() < deadline:
|
||||
detail = self.manager.detail(session["id"])
|
||||
if (
|
||||
detail["state"] == "awaiting_approval"
|
||||
and detail["proposals"][-1]["id"] != old_proposal["id"]
|
||||
):
|
||||
break
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
self.fail("rollback base did not produce a replacement proposal")
|
||||
self.assertEqual(detail["proposals"][-1]["baseTrialId"], baseline["id"])
|
||||
self.manager.cancel(session["id"])
|
||||
self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -18,10 +18,12 @@ from .process import GpuLease, ResourceBusyError, terminate_process
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
merge_proposal,
|
||||
validate_configuration_constraints,
|
||||
validate_constraints,
|
||||
validate_proposal,
|
||||
)
|
||||
from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation, validate_objective_weights
|
||||
from .storage import TuningStorage, now_iso
|
||||
from .storage import StorageConflict, TuningStorage, now_iso
|
||||
from .study import OptunaStudies
|
||||
from .tensorboard import ingest_scalars
|
||||
|
||||
@@ -98,7 +100,7 @@ class TuningManager:
|
||||
)
|
||||
):
|
||||
raise TuningError("gpuIds 必须是非空非负整数数组")
|
||||
trial_count = self._integer(payload, "trialCount", 12, 4, 20)
|
||||
trial_count = self._integer(payload, "trialCount", 12, 1, 100)
|
||||
rung0 = self._integer(payload, "initialIterations", 300, 1, 1000000)
|
||||
rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000)
|
||||
rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000)
|
||||
@@ -159,6 +161,15 @@ class TuningManager:
|
||||
session["trials"] = self.storage.list_trials(session_id)
|
||||
session["proposals"] = self.storage.list_proposals(session_id)
|
||||
session["audit"] = self.storage.audit_events(session_id)
|
||||
control = self.storage.get_control(session_id)
|
||||
current = next(
|
||||
(trial for trial in session["trials"] if trial["id"] == session["currentTrialId"]),
|
||||
None,
|
||||
)
|
||||
control["effectiveAfterCurrent"] = bool(
|
||||
current and current["state"] in {"training", "evaluating"}
|
||||
)
|
||||
session["control"] = control
|
||||
return session
|
||||
|
||||
def list(self) -> list[dict]:
|
||||
@@ -274,8 +285,8 @@ class TuningManager:
|
||||
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])
|
||||
|
||||
# 暂停只关闭后续 Trial 调度门;已经开始的 Trial 必须连同固定评估一起
|
||||
# 完成,避免把一次单步令牌错误地消耗在半个 Trial 上。
|
||||
eval_output = run_dir / "evaluation.json"
|
||||
eval_command = [
|
||||
self.python,
|
||||
@@ -332,6 +343,9 @@ class TuningManager:
|
||||
score=scored["score"],
|
||||
eligible=scored["eligible"],
|
||||
)
|
||||
best = self._best_highest_rung(session_id)
|
||||
if best is not None:
|
||||
self.storage.update_session(session_id, best_trial_id=best["id"])
|
||||
try:
|
||||
optuna_number = self.studies.record(
|
||||
session_id,
|
||||
@@ -363,6 +377,13 @@ class TuningManager:
|
||||
default=None,
|
||||
)
|
||||
|
||||
def _best_highest_rung(self, session_id: str) -> dict | None:
|
||||
for rung in (2, 1, 0):
|
||||
best = self._best(session_id, rung=rung)
|
||||
if best is not None:
|
||||
return best
|
||||
return None
|
||||
|
||||
def _proposal_context(self, session: dict) -> dict:
|
||||
trials = self.storage.list_trials(session["id"])[-12:]
|
||||
rejected_feedback = [
|
||||
@@ -374,6 +395,7 @@ class TuningManager:
|
||||
"task": session["config"]["taskId"],
|
||||
"objectiveWeights": session["objectiveWeights"],
|
||||
"allowlist": "服务端将验证固定 schema;最多四项修改",
|
||||
"parameterConstraints": self.storage.get_control(session["id"])["constraints"],
|
||||
"rejectedFeedback": rejected_feedback,
|
||||
"trials": [
|
||||
{
|
||||
@@ -414,6 +436,8 @@ class TuningManager:
|
||||
"model": "optuna-fallback",
|
||||
}
|
||||
source = "fallback"
|
||||
constraints = self.storage.get_control(session["id"])["constraints"]
|
||||
result["patch"] = validate_proposal(result["patch"], previous, constraints)
|
||||
proposal = self.storage.create_proposal(
|
||||
session["id"],
|
||||
base_trial_id,
|
||||
@@ -487,7 +511,10 @@ class TuningManager:
|
||||
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)
|
||||
best_score = max(
|
||||
(t["score"] if t["score"] is not None else float("-inf") for t in completed_rung0),
|
||||
default=0.0,
|
||||
)
|
||||
no_improve = 0
|
||||
while (
|
||||
next_number < session["config"]["trialCount"]
|
||||
@@ -495,12 +522,19 @@ class TuningManager:
|
||||
):
|
||||
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]
|
||||
self._wait_for_dispatch(session_id, cancel)
|
||||
control = self.storage.get_control(session_id)
|
||||
base = (
|
||||
self.storage.get_trial(control["activeBaseTrialId"])
|
||||
if control["activeBaseTrialId"]
|
||||
else 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":
|
||||
# 模式允许在 session 运行期间切换,因此每次决策都读取最新持久化值。
|
||||
current_mode = self.storage.get_session(session_id)["mode"]
|
||||
if current_mode == "approval":
|
||||
self.storage.update_session(
|
||||
session_id, state="awaiting_approval", message="等待批准 Agent 建议"
|
||||
)
|
||||
@@ -515,8 +549,9 @@ class TuningManager:
|
||||
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"])
|
||||
self._claim_dispatch(session_id, cancel)
|
||||
constraints = self.storage.get_control(session_id)["constraints"]
|
||||
reward_config = merge_proposal(base["rewardConfig"], proposal["patch"], constraints)
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
next_number,
|
||||
@@ -527,8 +562,16 @@ class TuningManager:
|
||||
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"]
|
||||
if control["activeBaseTrialId"]:
|
||||
self.storage.set_active_base(session_id, None)
|
||||
self._pause_after_step(session_id)
|
||||
result_score = result["score"]
|
||||
if (
|
||||
result["eligible"]
|
||||
and result_score is not None
|
||||
and result_score > best_score + 0.01
|
||||
):
|
||||
best_score = result_score
|
||||
no_improve = 0
|
||||
else:
|
||||
no_improve += 1
|
||||
@@ -537,13 +580,17 @@ class TuningManager:
|
||||
|
||||
# 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)
|
||||
previous.sort(
|
||||
key=lambda trial: (
|
||||
trial["score"] if trial["score"] is not None else float("-inf")
|
||||
),
|
||||
reverse=True,
|
||||
)
|
||||
promoted_numbers = {
|
||||
trial["number"]
|
||||
for trial in self.storage.list_trials(session_id)
|
||||
@@ -554,6 +601,7 @@ class TuningManager:
|
||||
continue
|
||||
if cancel.is_set():
|
||||
raise TuningError("session 已取消")
|
||||
self._claim_dispatch(session_id, cancel)
|
||||
root = self._session_root(session_id)
|
||||
checkpoint = root / parent["checkpointPath"]
|
||||
trial = self.storage.create_trial(
|
||||
@@ -566,12 +614,9 @@ class TuningManager:
|
||||
f"trial-{parent['number']:03d}-rung-{rung}",
|
||||
)
|
||||
self._execute_trial(session, trial, checkpoint)
|
||||
self._pause_after_step(session_id)
|
||||
|
||||
best = (
|
||||
self._best(session_id, rung=2)
|
||||
or self._best(session_id, rung=1)
|
||||
or self._best(session_id, rung=0)
|
||||
)
|
||||
best = self._best_highest_rung(session_id)
|
||||
if best is None:
|
||||
raise TuningError("没有通过安全门槛的 trial")
|
||||
preset_name = f"{session['config']['runName']}-{session_id[:8]}"
|
||||
@@ -602,10 +647,42 @@ class TuningManager:
|
||||
self.workers.pop(session_id, None)
|
||||
self.processes.pop(session_id, None)
|
||||
|
||||
def _wait_if_paused(self, session_id: str, cancel: threading.Event) -> None:
|
||||
def _wait_for_dispatch(self, session_id: str, cancel: threading.Event) -> None:
|
||||
"""Wait at a scheduler boundary without consuming a one-shot token."""
|
||||
with self.condition:
|
||||
while self.storage.get_session(session_id)["state"] == "paused" and not cancel.is_set():
|
||||
while not cancel.is_set():
|
||||
session = self.storage.get_session(session_id)
|
||||
control = self.storage.get_control(session_id)
|
||||
if session["state"] == "paused":
|
||||
self.condition.wait(timeout=1.0)
|
||||
continue
|
||||
if control["runPolicy"] == "continuous" or control["dispatchTokens"] > 0:
|
||||
return
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="paused",
|
||||
message="单步 Trial 已完成;等待下一个调度令牌",
|
||||
)
|
||||
self.storage.audit(session_id, "step_gate_waiting", {})
|
||||
self.condition.wait(timeout=1.0)
|
||||
raise TuningError("session 已取消")
|
||||
|
||||
def _claim_dispatch(self, session_id: str, cancel: threading.Event) -> None:
|
||||
while not cancel.is_set():
|
||||
self._wait_for_dispatch(session_id, cancel)
|
||||
if self.storage.use_dispatch_token(session_id):
|
||||
return
|
||||
raise TuningError("session 已取消")
|
||||
|
||||
def _pause_after_step(self, session_id: str) -> None:
|
||||
control = self.storage.get_control(session_id)
|
||||
if control["runPolicy"] == "step" and control["dispatchTokens"] == 0:
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="paused",
|
||||
message="单步 Trial 已完成;后续调度已暂停",
|
||||
)
|
||||
self.storage.audit(session_id, "step_trial_completed", {})
|
||||
|
||||
def approve(self, session_id: str, proposal_id: str, payload: Any) -> dict:
|
||||
proposal = self.storage.get_proposal(proposal_id)
|
||||
@@ -613,11 +690,15 @@ class TuningManager:
|
||||
raise TuningError("proposal 不属于该 session")
|
||||
patch = proposal["patch"]
|
||||
feedback = None
|
||||
base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
constraints = self.storage.get_control(session_id)["constraints"]
|
||||
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"])
|
||||
patch = payload["patch"]
|
||||
if feedback is not None and (not isinstance(feedback, str) or len(feedback) > 2000):
|
||||
raise TuningError("feedback 无效")
|
||||
patch = validate_proposal(patch, base["rewardConfig"], constraints)
|
||||
if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch):
|
||||
raise TuningError("proposal 已处理")
|
||||
self.storage.audit(
|
||||
@@ -642,22 +723,213 @@ class TuningManager:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def set_mode(self, session_id: str, payload: Any) -> dict:
|
||||
if not isinstance(payload, dict) or payload.get("mode") not in ("automatic", "approval"):
|
||||
raise TuningError("mode 必须是 automatic 或 approval")
|
||||
mode = payload["mode"]
|
||||
with self.condition:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in RUNNING_STATES | {"awaiting_approval", "paused"}:
|
||||
raise TuningError("当前状态不能切换运行模式")
|
||||
previous = session["mode"]
|
||||
if previous == mode:
|
||||
return self.detail(session_id)
|
||||
self.storage.update_session(session_id, mode=mode)
|
||||
approved_ids = []
|
||||
if mode == "automatic":
|
||||
constraints = self.storage.get_control(session_id)["constraints"]
|
||||
for proposal in self.storage.list_proposals(session_id):
|
||||
if proposal["state"] != "pending":
|
||||
continue
|
||||
base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
try:
|
||||
validate_proposal(proposal["patch"], base["rewardConfig"], constraints)
|
||||
except Exception as error:
|
||||
self.storage.decide_proposal(
|
||||
proposal["id"], "rejected", f"参数护栏已变化:{error}"
|
||||
)
|
||||
continue
|
||||
if self.storage.decide_proposal(
|
||||
proposal["id"], "approved", "运行时切换为全自动模式"
|
||||
):
|
||||
approved_ids.append(proposal["id"])
|
||||
if session["state"] == "awaiting_approval":
|
||||
self.storage.update_session(
|
||||
session_id, state="running", message="已切换为全自动模式,继续调参"
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"session_mode_changed",
|
||||
{"from": previous, "to": mode, "autoApprovedProposalIds": approved_ids},
|
||||
)
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def set_constraints(self, session_id: str, payload: Any) -> dict:
|
||||
if not isinstance(payload, dict) or set(payload) != {"revision", "constraints"}:
|
||||
raise TuningError("参数护栏请求必须包含 revision 与 constraints")
|
||||
revision = payload["revision"]
|
||||
if isinstance(revision, bool) or not isinstance(revision, int) or revision < 0:
|
||||
raise TuningError("constraints revision 必须是非负整数")
|
||||
constraints = validate_constraints(payload["constraints"])
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in ACTIVE_SESSION_STATES:
|
||||
raise TuningError("终态 session 不能修改参数护栏")
|
||||
control = self.storage.get_control(session_id)
|
||||
base = None
|
||||
if control["activeBaseTrialId"]:
|
||||
base = self.storage.get_trial(control["activeBaseTrialId"])
|
||||
elif session["currentTrialId"]:
|
||||
base = self.storage.get_trial(session["currentTrialId"])
|
||||
else:
|
||||
base = self._best_highest_rung(session_id)
|
||||
if base is not None:
|
||||
validate_configuration_constraints(base["rewardConfig"], constraints)
|
||||
try:
|
||||
updated = self.storage.replace_constraints(session_id, revision, constraints)
|
||||
except StorageConflict as error:
|
||||
raise ResourceBusyError(str(error)) from error
|
||||
|
||||
rejected = []
|
||||
for proposal in self.storage.list_proposals(session_id):
|
||||
if proposal["state"] != "pending":
|
||||
continue
|
||||
proposal_base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
try:
|
||||
validate_proposal(proposal["patch"], proposal_base["rewardConfig"], constraints)
|
||||
except Exception as error:
|
||||
message = f"参数护栏 revision {updated['constraintsRevision']}:{error}"
|
||||
if self.storage.decide_proposal(proposal["id"], "rejected", message):
|
||||
rejected.append(proposal["id"])
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"constraints_updated",
|
||||
{
|
||||
"revision": updated["constraintsRevision"],
|
||||
"paths": sorted(constraints),
|
||||
"rejectedProposalIds": rejected,
|
||||
},
|
||||
)
|
||||
if rejected:
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def step(self, session_id: str, payload: Any) -> dict:
|
||||
if not isinstance(payload, dict) or payload.get("count", 1) != 1:
|
||||
raise TuningError("单步调度一次只能发放 1 个 Trial 令牌")
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in {"paused", "awaiting_approval"}:
|
||||
raise TuningError("请先暂停或等待 Proposal 审批,再执行单步 Trial")
|
||||
control = self.storage.get_control(session_id)
|
||||
if control["dispatchTokens"] > 0:
|
||||
raise TuningError("已有未消费的单步 Trial 令牌")
|
||||
try:
|
||||
self.storage.grant_dispatch_token(session_id)
|
||||
except StorageConflict as error:
|
||||
raise ResourceBusyError(str(error)) from error
|
||||
if session["state"] == "paused":
|
||||
has_pending = any(
|
||||
proposal["state"] == "pending"
|
||||
for proposal in self.storage.list_proposals(session_id)
|
||||
)
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="awaiting_approval" if has_pending else "running",
|
||||
message="单步令牌已就绪;请审批 Proposal"
|
||||
if has_pending
|
||||
else "已授权执行一个 Trial",
|
||||
)
|
||||
self.storage.audit(session_id, "step_token_granted", {"count": 1})
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
return self.detail(session_id)
|
||||
|
||||
def rollback(self, session_id: str, payload: Any) -> dict:
|
||||
if not isinstance(payload, dict):
|
||||
raise TuningError("rollback 请求体必须是对象")
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in {"paused", "awaiting_approval"}:
|
||||
raise TuningError("回滚只能在安全暂停或等待审批时执行")
|
||||
target_id = payload.get("trialId")
|
||||
if payload.get("target") == "best":
|
||||
target_id = session["bestTrialId"] or (self._best_highest_rung(session_id) or {}).get(
|
||||
"id"
|
||||
)
|
||||
if not isinstance(target_id, str):
|
||||
raise TuningError("rollback 必须指定 trialId 或 target=best")
|
||||
target = self.storage.get_trial(target_id)
|
||||
if target["sessionId"] != session_id:
|
||||
raise TuningError("rollback Trial 不属于该 session")
|
||||
if target["state"] != "completed" or not target["eligible"]:
|
||||
raise TuningError("只能回滚到已完成且通过安全门槛的 Trial")
|
||||
checkpoint = payload.get("checkpoint", False)
|
||||
if not isinstance(checkpoint, bool):
|
||||
raise TuningError("checkpoint 必须是布尔值")
|
||||
if checkpoint:
|
||||
path_value = target.get("checkpointPath")
|
||||
if not path_value:
|
||||
raise TuningError("目标 Trial 没有可用 checkpoint")
|
||||
root = self._session_root(session_id)
|
||||
path = (root / path_value).resolve()
|
||||
if not path.is_relative_to(root) or not path.is_file():
|
||||
raise TuningError("目标 checkpoint 不存在或路径非法")
|
||||
|
||||
superseded = []
|
||||
for proposal in self.storage.list_proposals(session_id):
|
||||
if proposal["state"] == "pending" and self.storage.decide_proposal(
|
||||
proposal["id"], "rejected", f"由回滚到 Trial {target_id[:8]} 取代"
|
||||
):
|
||||
superseded.append(proposal["id"])
|
||||
self.storage.set_active_base(session_id, target_id)
|
||||
self.storage.reset_step_gate(session_id)
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="paused",
|
||||
message=f"已回滚到 Trial {target['number']} / Rung {target['rung']};等待单步或继续",
|
||||
)
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"rollback_selected",
|
||||
{
|
||||
"trialId": target_id,
|
||||
"checkpoint": checkpoint,
|
||||
"supersededProposalIds": superseded,
|
||||
},
|
||||
)
|
||||
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("当前状态不能暂停")
|
||||
# 以持久化调度门记录暂停意图,而不是只依赖 session.state。当前 Trial
|
||||
# 的训练/评估会继续完成,下一次 _claim_dispatch 必须等待显式继续或单步。
|
||||
self.storage.reset_step_gate(session_id)
|
||||
self.storage.update_session(
|
||||
session_id, state="paused", message="已暂停后续调度;当前子进程将完成"
|
||||
session_id, state="paused", message="已暂停后续调度;在途 Trial(如有)将完整结束"
|
||||
)
|
||||
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="继续调参")
|
||||
self.storage.set_run_policy(session_id, "continuous")
|
||||
has_pending = any(
|
||||
proposal["state"] == "pending"
|
||||
for proposal in self.storage.list_proposals(session_id)
|
||||
)
|
||||
self.storage.update_session(
|
||||
session_id,
|
||||
state="awaiting_approval" if has_pending else "running",
|
||||
message="请审批待处理 Proposal" if has_pending else "连续调参已恢复",
|
||||
)
|
||||
with self.condition:
|
||||
self.condition.notify_all()
|
||||
elif session["state"] == "interrupted":
|
||||
self.storage.set_run_policy(session_id, "continuous")
|
||||
self.storage.update_session(session_id, state="queued", message="从最近完整结果恢复")
|
||||
self._start_worker(session_id, resume=True)
|
||||
else:
|
||||
@@ -665,7 +937,9 @@ class TuningManager:
|
||||
return self.detail(session_id)
|
||||
|
||||
def cancel(self, session_id: str) -> dict:
|
||||
self.storage.get_session(session_id)
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["state"] not in ACTIVE_SESSION_STATES:
|
||||
return self.detail(session_id)
|
||||
self.storage.update_session(session_id, state="cancelled", message="正在取消")
|
||||
event = self.cancel_events.get(session_id)
|
||||
if event:
|
||||
@@ -678,12 +952,22 @@ class TuningManager:
|
||||
return self.detail(session_id)
|
||||
|
||||
def metrics(
|
||||
self, session_id: str, trial_id: str, tags: list[str] | None, max_points: int
|
||||
self,
|
||||
session_id: str,
|
||||
trial_id: str,
|
||||
tags: list[str] | None,
|
||||
max_points: int,
|
||||
after_step: int | None = None,
|
||||
) -> 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)}
|
||||
series = self.storage.metrics(trial_id, tags, max_points, after_step)
|
||||
next_step = max(
|
||||
(point["step"] for item in series for point in item["points"]),
|
||||
default=after_step,
|
||||
)
|
||||
return {"trialId": trial_id, "series": series, "nextStep": next_step}
|
||||
|
||||
def best_artifact(self, session_id: str) -> Path:
|
||||
session = self.storage.get_session(session_id)
|
||||
|
||||
@@ -93,6 +93,68 @@ def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None:
|
||||
raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold")
|
||||
|
||||
|
||||
def _path_spec(path: str) -> tuple[str, str, NumericSpec]:
|
||||
if path.startswith("weights."):
|
||||
section, name = "weights", path.removeprefix("weights.")
|
||||
spec = WEIGHT_SPECS.get(name)
|
||||
elif path.startswith("params."):
|
||||
section, name = "params", path.removeprefix("params.")
|
||||
spec = PARAMETER_SPECS.get(name)
|
||||
else:
|
||||
section, name, spec = "", "", None
|
||||
if spec is None:
|
||||
raise RewardConfigError(f"未知参数约束:{path}")
|
||||
return section, name, spec
|
||||
|
||||
|
||||
def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]:
|
||||
"""Validate sparse per-session range/fixed safety constraints."""
|
||||
root = _mapping(value, "constraints")
|
||||
if len(root) > len(WEIGHT_SPECS) + len(PARAMETER_SPECS):
|
||||
raise RewardConfigError("constraints 数量超过白名单参数总数")
|
||||
result: dict[str, dict[str, float | str]] = {}
|
||||
for raw_path, raw_constraint in root.items():
|
||||
if not isinstance(raw_path, str):
|
||||
raise RewardConfigError("constraint path 必须是字符串")
|
||||
_, _, spec = _path_spec(raw_path)
|
||||
constraint = _mapping(raw_constraint, raw_path)
|
||||
kind = constraint.get("kind")
|
||||
if kind == "fixed":
|
||||
if set(constraint) != {"kind", "value"}:
|
||||
raise RewardConfigError(f"{raw_path} fixed 约束只能包含 kind/value")
|
||||
fixed = _number(f"{raw_path}.value", constraint["value"], spec)
|
||||
result[raw_path] = {"kind": "fixed", "value": fixed}
|
||||
elif kind == "range":
|
||||
if set(constraint) != {"kind", "min", "max"}:
|
||||
raise RewardConfigError(f"{raw_path} range 约束只能包含 kind/min/max")
|
||||
minimum = _number(f"{raw_path}.min", constraint["min"], spec)
|
||||
maximum = _number(f"{raw_path}.max", constraint["max"], spec)
|
||||
if minimum > maximum:
|
||||
raise RewardConfigError(f"{raw_path} 下限不能大于上限")
|
||||
result[raw_path] = {"kind": "range", "min": minimum, "max": maximum}
|
||||
else:
|
||||
raise RewardConfigError(f"{raw_path}.kind 必须是 range 或 fixed")
|
||||
return result
|
||||
|
||||
|
||||
def validate_configuration_constraints(value: Any, constraints: Any) -> None:
|
||||
"""Ensure a complete reward configuration satisfies every session constraint."""
|
||||
config = validate_configuration(value)
|
||||
checked = validate_constraints(constraints)
|
||||
for path, constraint in checked.items():
|
||||
section, name, _ = _path_spec(path)
|
||||
current = config[section][name]
|
||||
if constraint["kind"] == "fixed":
|
||||
if current != constraint["value"]:
|
||||
raise RewardConfigError(
|
||||
f"{path} 已固定为 {constraint['value']},不能设为 {current}"
|
||||
)
|
||||
elif current < constraint["min"] or current > constraint["max"]:
|
||||
raise RewardConfigError(
|
||||
f"{path}={current} 超出工程锁定范围 {constraint['min']}–{constraint['max']}"
|
||||
)
|
||||
|
||||
|
||||
def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a complete configuration and reject missing/unknown fields."""
|
||||
root = _mapping(value, "rewardConfig")
|
||||
@@ -118,8 +180,10 @@ def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
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."""
|
||||
def validate_proposal(
|
||||
value: Any, previous: Any, constraints: Any | None = None
|
||||
) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config and guardrails."""
|
||||
current = validate_configuration(previous)
|
||||
root = _mapping(value, "proposal")
|
||||
if not set(root).issubset({"weights", "params"}):
|
||||
@@ -167,12 +231,16 @@ def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]:
|
||||
candidate["weights"].update(patch["weights"])
|
||||
candidate["params"].update(patch["params"])
|
||||
_cross_validate(candidate)
|
||||
if constraints is not None:
|
||||
validate_configuration_constraints(candidate, constraints)
|
||||
return patch
|
||||
|
||||
|
||||
def merge_proposal(previous: Any, proposal: Any) -> dict[str, dict[str, float]]:
|
||||
def merge_proposal(
|
||||
previous: Any, proposal: Any, constraints: Any | None = None
|
||||
) -> dict[str, dict[str, float]]:
|
||||
current = validate_configuration(previous)
|
||||
patch = validate_proposal(proposal, current)
|
||||
patch = validate_proposal(proposal, current, constraints)
|
||||
merged = deepcopy(current)
|
||||
merged["weights"].update(patch["weights"])
|
||||
merged["params"].update(patch["params"])
|
||||
|
||||
@@ -12,7 +12,11 @@ from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
SCHEMA_VERSION = 2
|
||||
|
||||
|
||||
class StorageConflict(RuntimeError):
|
||||
"""Optimistic-concurrency revision mismatch."""
|
||||
|
||||
|
||||
def now_iso() -> str:
|
||||
@@ -135,11 +139,29 @@ class TuningStorage:
|
||||
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 TABLE IF NOT EXISTS session_controls(
|
||||
session_id TEXT PRIMARY KEY REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
run_policy TEXT NOT NULL DEFAULT 'continuous'
|
||||
CHECK(run_policy IN ('continuous','step')),
|
||||
dispatch_tokens INTEGER NOT NULL DEFAULT 0 CHECK(dispatch_tokens >= 0),
|
||||
active_base_trial_id TEXT,
|
||||
revision INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS session_constraints(
|
||||
session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
path TEXT NOT NULL, kind TEXT NOT NULL CHECK(kind IN ('range','fixed')),
|
||||
min_value REAL, max_value REAL, fixed_value REAL,
|
||||
updated_at TEXT NOT NULL,
|
||||
PRIMARY KEY(session_id, path)
|
||||
);
|
||||
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 session_controls(session_id) SELECT id FROM sessions"
|
||||
)
|
||||
connection.execute(
|
||||
"INSERT OR IGNORE INTO schema_migrations(version) VALUES (?)", (SCHEMA_VERSION,)
|
||||
)
|
||||
@@ -170,6 +192,7 @@ class TuningStorage:
|
||||
"VALUES (?, 'queued', ?, ?, ?, ?, ?, '等待基线训练', ?)",
|
||||
(session_id, mode, at, at, _json(config), _json(objective), int(fallback)),
|
||||
)
|
||||
connection.execute("INSERT INTO session_controls(session_id) VALUES (?)", (session_id,))
|
||||
connection.execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
@@ -212,6 +235,7 @@ class TuningStorage:
|
||||
def update_session(self, session_id: str, **changes: Any) -> bool:
|
||||
columns = {
|
||||
"state": "state",
|
||||
"mode": "mode",
|
||||
"message": "message",
|
||||
"current_trial_id": "current_trial_id",
|
||||
"best_trial_id": "best_trial_id",
|
||||
@@ -230,6 +254,143 @@ class TuningStorage:
|
||||
)
|
||||
return cursor.rowcount == 1
|
||||
|
||||
def get_control(self, session_id: str) -> dict:
|
||||
row = (
|
||||
self.connection()
|
||||
.execute("SELECT * FROM session_controls WHERE session_id=?", (session_id,))
|
||||
.fetchone()
|
||||
)
|
||||
if row is None:
|
||||
raise KeyError(session_id)
|
||||
constraint_rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
"SELECT * FROM session_constraints WHERE session_id=? ORDER BY path", (session_id,)
|
||||
)
|
||||
.fetchall()
|
||||
)
|
||||
constraints = {}
|
||||
for constraint in constraint_rows:
|
||||
if constraint["kind"] == "fixed":
|
||||
value = {"kind": "fixed", "value": constraint["fixed_value"]}
|
||||
else:
|
||||
value = {
|
||||
"kind": "range",
|
||||
"min": constraint["min_value"],
|
||||
"max": constraint["max_value"],
|
||||
}
|
||||
constraints[constraint["path"]] = value
|
||||
return {
|
||||
"runPolicy": row["run_policy"],
|
||||
"dispatchTokens": row["dispatch_tokens"],
|
||||
"constraintsRevision": row["revision"],
|
||||
"constraints": constraints,
|
||||
"activeBaseTrialId": row["active_base_trial_id"],
|
||||
}
|
||||
|
||||
def replace_constraints(
|
||||
self, session_id: str, expected_revision: int, constraints: dict
|
||||
) -> dict:
|
||||
at = now_iso()
|
||||
with self.transaction() as connection:
|
||||
row = connection.execute(
|
||||
"SELECT revision FROM session_controls WHERE session_id=?", (session_id,)
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(session_id)
|
||||
if row["revision"] != expected_revision:
|
||||
raise StorageConflict(
|
||||
f"参数护栏 revision 已变化(当前 {row['revision']},请求 {expected_revision})"
|
||||
)
|
||||
connection.execute("DELETE FROM session_constraints WHERE session_id=?", (session_id,))
|
||||
for path, constraint in constraints.items():
|
||||
connection.execute(
|
||||
"INSERT INTO session_constraints("
|
||||
"session_id,path,kind,min_value,max_value,fixed_value,updated_at) "
|
||||
"VALUES (?,?,?,?,?,?,?)",
|
||||
(
|
||||
session_id,
|
||||
path,
|
||||
constraint["kind"],
|
||||
constraint.get("min"),
|
||||
constraint.get("max"),
|
||||
constraint.get("value"),
|
||||
at,
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE session_controls SET revision=revision+1 WHERE session_id=?",
|
||||
(session_id,),
|
||||
)
|
||||
return self.get_control(session_id)
|
||||
|
||||
def grant_dispatch_token(self, session_id: str) -> dict:
|
||||
"""Atomically grant the sole outstanding one-Trial token."""
|
||||
with self.transaction() as connection:
|
||||
cursor = connection.execute(
|
||||
"UPDATE session_controls SET run_policy='step',dispatch_tokens=1 "
|
||||
"WHERE session_id=? AND dispatch_tokens=0",
|
||||
(session_id,),
|
||||
)
|
||||
if cursor.rowcount != 1:
|
||||
exists = connection.execute(
|
||||
"SELECT 1 FROM session_controls WHERE session_id=?", (session_id,)
|
||||
).fetchone()
|
||||
if exists is None:
|
||||
raise KeyError(session_id)
|
||||
raise StorageConflict("已有未消费的单步 Trial 令牌")
|
||||
return self.get_control(session_id)
|
||||
|
||||
def use_dispatch_token(self, session_id: str) -> bool:
|
||||
"""Atomically consume one step token; continuous mode never needs a token."""
|
||||
with self.transaction() as connection:
|
||||
row = connection.execute(
|
||||
"SELECT run_policy,dispatch_tokens FROM session_controls WHERE session_id=?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(session_id)
|
||||
if row["run_policy"] == "continuous":
|
||||
return True
|
||||
if row["dispatch_tokens"] <= 0:
|
||||
return False
|
||||
connection.execute(
|
||||
"UPDATE session_controls SET dispatch_tokens=dispatch_tokens-1 WHERE session_id=?",
|
||||
(session_id,),
|
||||
)
|
||||
return True
|
||||
|
||||
def set_run_policy(self, session_id: str, policy: str) -> dict:
|
||||
if policy not in {"continuous", "step"}:
|
||||
raise ValueError(policy)
|
||||
cursor = self.connection().execute(
|
||||
"UPDATE session_controls SET run_policy=?,"
|
||||
"dispatch_tokens=CASE WHEN ?='continuous' THEN 0 ELSE dispatch_tokens END "
|
||||
"WHERE session_id=?",
|
||||
(policy, policy, session_id),
|
||||
)
|
||||
if cursor.rowcount != 1:
|
||||
raise KeyError(session_id)
|
||||
return self.get_control(session_id)
|
||||
|
||||
def reset_step_gate(self, session_id: str) -> dict:
|
||||
cursor = self.connection().execute(
|
||||
"UPDATE session_controls SET run_policy='step',dispatch_tokens=0 WHERE session_id=?",
|
||||
(session_id,),
|
||||
)
|
||||
if cursor.rowcount != 1:
|
||||
raise KeyError(session_id)
|
||||
return self.get_control(session_id)
|
||||
|
||||
def set_active_base(self, session_id: str, trial_id: str | None) -> dict:
|
||||
cursor = self.connection().execute(
|
||||
"UPDATE session_controls SET active_base_trial_id=? WHERE session_id=?",
|
||||
(trial_id, session_id),
|
||||
)
|
||||
if cursor.rowcount != 1:
|
||||
raise KeyError(session_id)
|
||||
return self.get_control(session_id)
|
||||
|
||||
def create_trial(
|
||||
self,
|
||||
session_id: str,
|
||||
@@ -424,13 +585,20 @@ class TuningStorage:
|
||||
)
|
||||
|
||||
def metrics(
|
||||
self, trial_id: str, tags: list[str] | None = None, max_points: int = 1000
|
||||
self,
|
||||
trial_id: str,
|
||||
tags: list[str] | None = None,
|
||||
max_points: int = 1000,
|
||||
after_step: int | None = None,
|
||||
) -> list[dict]:
|
||||
parameters: list[Any] = [trial_id]
|
||||
clause = "trial_id=?"
|
||||
if tags:
|
||||
clause += f" AND tag IN ({','.join('?' for _ in tags)})"
|
||||
parameters.extend(tags)
|
||||
if after_step is not None:
|
||||
clause += " AND step>?"
|
||||
parameters.append(after_step)
|
||||
rows = (
|
||||
self.connection()
|
||||
.execute(
|
||||
|
||||
Reference in New Issue
Block a user