feat(tuning): release V0.8.2 Agent 界面重构
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-03 16:25:53 +08:00
parent a9b07e0abf
commit 63d67a645b
35 changed files with 4847 additions and 811 deletions
+7 -3
View File
@@ -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×` 变化率校验。
## 测试
+41 -4
View File
@@ -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()
+45 -1
View File
@@ -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()
+312 -28
View File
@@ -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)
+72 -4
View File
@@ -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"])
+170 -2
View File
@@ -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(