feat(tuning): release V0.8.2 Agent 界面重构
This commit is contained in:
@@ -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