"""Persistent reward tuning session orchestrator.""" from __future__ import annotations import json import os import re import shutil import subprocess import threading from contextlib import suppress from copy import deepcopy from pathlib import Path from typing import Any from .advisor import AdvisorUnavailable, DeepSeekAdvisor from .process import GpuLease, ResourceBusyError, terminate_process from .schema import ( BASE_REWARD_CONFIGURATION, merge_proposal, validate_configuration_constraints, validate_constraints, validate_proposal, ) from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation, validate_objective_weights from .storage import StorageConflict, TuningStorage, now_iso from .study import OptunaStudies from .tensorboard import ingest_scalars ACTIVE_SESSION_STATES = { "queued", "running", "evaluating", "awaiting_approval", "paused", "interrupted", } RUNNING_STATES = {"queued", "running", "evaluating"} RUN_NAME = re.compile(r"^[A-Za-z0-9_.-]{1,64}$") class TuningError(RuntimeError): pass class TuningManager: def __init__( self, trainer_root: Path, python: str, data_root: Path, lease: GpuLease, advisor: Any | None = None, ): self.trainer_root = trainer_root.expanduser().resolve() self.python = python self.data_root = data_root.expanduser().resolve() self.data_root.mkdir(parents=True, exist_ok=True) self.storage = TuningStorage(self.data_root / "tuning.sqlite3") self.storage.recover_interrupted() self.studies = OptunaStudies(self.data_root) self.lease = lease self.advisor = advisor or DeepSeekAdvisor() self.lock = threading.RLock() self.condition = threading.Condition(self.lock) self.workers: dict[str, threading.Thread] = {} self.processes: dict[str, subprocess.Popen[str]] = {} self.cancel_events: dict[str, threading.Event] = {} def capability(self) -> dict[str, Any]: capability = self.advisor.capability() capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()}) return capability @staticmethod def _integer(payload: dict, name: str, default: int, minimum: int, maximum: int) -> int: value = payload.get(name, default) if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum: raise TuningError(f"{name} 必须在 {minimum}–{maximum} 之间") return value def parse_create(self, payload: Any) -> tuple[str, dict, dict, bool]: if not isinstance(payload, dict): raise TuningError("请求体必须是 JSON 对象") mode = payload.get("mode", "automatic") if mode not in ("automatic", "approval"): raise TuningError("mode 必须是 automatic 或 approval") if payload.get("taskId", "Unitree-Go2-Flat") != "Unitree-Go2-Flat": raise TuningError("第一版只支持 Unitree-Go2-Flat") run_name = payload.get("runName", "auto-tune") if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name): raise TuningError("runName 格式无效") gpu_ids = payload.get("gpuIds", [0]) if ( not isinstance(gpu_ids, list) or not gpu_ids or any( isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 255 for value in gpu_ids ) ): raise TuningError("gpuIds 必须是非空非负整数数组") trial_count = self._integer(payload, "trialCount", 12, 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) config = { "taskId": "Unitree-Go2-Flat", "numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384), "seed": self._integer(payload, "seed", 42, 0, 2147483647), "runName": run_name, "gpuIds": gpu_ids, "trialCount": trial_count, "rungs": [rung0, rung1, rung2], "promote": [trial_count, min(4, trial_count), min(2, trial_count)], "evalNumEnvs": self._integer(payload, "evalNumEnvs", 256, 1, 4096), "evalSteps": self._integer(payload, "evalSteps", 1000, 10, 100000), "earlyStopPatience": self._integer(payload, "earlyStopPatience", 4, 1, 20), } objective = validate_objective_weights( payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS) ) fallback = payload.get("fallbackEnabled", False) if not isinstance(fallback, bool): raise TuningError("fallbackEnabled 必须是布尔值") return mode, config, objective, fallback def create(self, payload: Any) -> dict: mode, config, objective, fallback = self.parse_create(payload) capability = self.capability() if not capability["ready"]: raise TuningError("评估入口未就绪") if not capability["configured"] and not fallback: raise TuningError("DeepSeek Agent 未配置;设置 DEEPSEEK_API_KEY 或显式启用 fallback") with self.lock: active = [ session for session in self.storage.list_sessions() if session["state"] in ACTIVE_SESSION_STATES ] if active: raise ResourceBusyError("已有调参 session 未结束") session = self.storage.create_session(mode, config, objective, fallback) self._start_worker(session["id"], resume=False) return self.detail(session["id"]) def _start_worker(self, session_id: str, resume: bool) -> None: cancel = threading.Event() self.cancel_events[session_id] = cancel worker = threading.Thread( target=self._run_session, args=(session_id, resume, cancel), name=f"tuning-{session_id[:8]}", daemon=True, ) self.workers[session_id] = worker worker.start() def detail(self, session_id: str) -> dict: session = self.storage.get_session(session_id) session["trials"] = self.storage.list_trials(session_id) session["proposals"] = self.storage.list_proposals(session_id) session["audit"] = self.storage.audit_events(session_id) 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]: return self.storage.list_sessions() def _session_root(self, session_id: str) -> Path: root = (self.data_root / "sessions" / session_id).resolve() if not root.is_relative_to(self.data_root): raise TuningError("非法 session 路径") root.mkdir(parents=True, exist_ok=True) return root def _run_command( self, session_id: str, command: list[str], cwd: Path, environment: dict[str, str], log_path: Path, ) -> int: owner = f"tuning:{session_id}" self.lease.acquire(owner) try: process = subprocess.Popen( command, cwd=cwd, env=environment, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, encoding="utf-8", errors="replace", bufsize=1, start_new_session=True, ) with self.lock: self.processes[session_id] = process assert process.stdout is not None with log_path.open("a", encoding="utf-8") as log: for line in process.stdout: log.write(line) log.flush() if self.cancel_events[session_id].is_set(): terminate_process(process) break return process.wait() finally: with self.lock: self.processes.pop(session_id, None) self.lease.release(owner) @staticmethod def _latest_checkpoint(run_dir: Path) -> Path | None: values = [] for path in run_dir.glob("model_*.pt"): match = re.fullmatch(r"model_(\d+)\.pt", path.name) if match: values.append((int(match.group(1)), path)) return max(values, default=(0, None), key=lambda value: value[0])[1] def _execute_trial( self, session: dict, trial: dict, resume_checkpoint: Path | None = None ) -> dict: session_id, trial_id = session["id"], trial["id"] root = self._session_root(session_id) run_dir = (root / trial["runDir"]).resolve() if not run_dir.is_relative_to(root): raise TuningError("trial 目录越界") run_dir.mkdir(parents=True, exist_ok=True) reward_path = run_dir / "reward_config.json" reward_path.write_text( json.dumps(trial["rewardConfig"], ensure_ascii=False, indent=2) + "\n", encoding="utf-8" ) config = session["config"] command = [ self.python, "-u", "scripts/train.py", config["taskId"], f"--env.scene.num-envs={config['numEnvs']}", f"--agent.max-iterations={trial['targetIterations']}", f"--agent.seed={config['seed']}", f"--agent.run-name={config['runName']}-t{trial['number']}-r{trial['rung']}", "--agent.logger=tensorboard", "--agent.upload-model=False", "--gpu-ids", json.dumps(config["gpuIds"], separators=(",", ":")), "--output-dir", str(run_dir), "--reward-config", str(reward_path), ] if resume_checkpoint is not None: command.extend(("--resume-checkpoint", str(resume_checkpoint))) environment = os.environ.copy() environment["WANDB_MODE"] = "disabled" environment["WANDB_SILENT"] = "true" self.storage.update_trial( trial_id, state="training", started_at=now_iso(), message="正在训练" ) self.storage.update_session( session_id, state="running", message=f"正在训练 trial {trial['number']} / rung {trial['rung']}", ) return_code = self._run_command( session_id, command, self.trainer_root, environment, run_dir / "train.log" ) ingest_scalars(self.storage, trial_id, run_dir) if self.cancel_events[session_id].is_set(): raise TuningError("session 已取消") checkpoint = self._latest_checkpoint(run_dir) policy = run_dir / "policy.onnx" if return_code != 0 or checkpoint is None or not policy.is_file(): raise TuningError(f"训练失败(返回码 {return_code})或缺少 checkpoint/policy.onnx") # 暂停只关闭后续 Trial 调度门;已经开始的 Trial 必须连同固定评估一起 # 完成,避免把一次单步令牌错误地消耗在半个 Trial 上。 eval_output = run_dir / "evaluation.json" eval_command = [ self.python, "-u", "scripts/evaluate.py", config["taskId"], "--checkpoint", str(checkpoint), "--output", str(eval_output), "--reward-config", str(reward_path), f"--num-envs={config['evalNumEnvs']}", f"--steps-per-seed={config['evalSteps']}", "--gpu-ids", json.dumps(config["gpuIds"], separators=(",", ":")), ] self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估") self.storage.update_session( session_id, state="evaluating", message=f"正在评估 trial {trial['number']}" ) return_code = self._run_command( session_id, eval_command, self.trainer_root, environment, run_dir / "evaluate.log" ) ingest_scalars(self.storage, trial_id, run_dir / "evaluation-events") if return_code != 0 or not eval_output.is_file(): raise TuningError(f"评估失败(返回码 {return_code})") evaluation = json.loads(eval_output.read_text(encoding="utf-8")) baseline_trial = self.storage.list_trials(session_id)[0] if baseline_trial["evaluation"] is None: scored = { "eligible": True, "score": 0.0, "components": {}, "metrics": evaluation["metrics"], } else: scored = score_evaluation( baseline_trial["evaluation"]["metrics"], evaluation["metrics"], session["objectiveWeights"], ) evaluation["score"] = scored rel_checkpoint = str(checkpoint.relative_to(root)) rel_policy = str(policy.relative_to(root)) self.storage.update_trial( trial_id, state="completed", ended_at=now_iso(), message="训练与评估完成", checkpoint_path=rel_checkpoint, policy_path=rel_policy, evaluation=evaluation, score=scored["score"], eligible=scored["eligible"], ) 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, trial["rewardConfig"], scored["score"], scored["eligible"], trial["rung"], ) self.storage.audit( session_id, "optuna_trial_recorded", {"trialId": trial_id, "number": optuna_number} ) except Exception as error: self.storage.audit( session_id, "optuna_record_failed", {"trialId": trial_id, "error": str(error)[:500]} ) return self.storage.get_trial(trial_id) def _best(self, session_id: str, rung: int | None = None) -> dict | None: trials = [ trial for trial in self.storage.list_trials(session_id) if trial["state"] == "completed" and trial["eligible"] ] if rung is not None: trials = [trial for trial in trials if trial["rung"] == rung] return max( trials, key=lambda trial: trial["score"] if trial["score"] is not None else -999, default=None, ) def _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 = [ proposal["feedback"] for proposal in self.storage.list_proposals(session["id"]) if proposal["state"] == "rejected" and proposal["feedback"] ][-4:] return { "task": session["config"]["taskId"], "objectiveWeights": session["objectiveWeights"], "allowlist": "服务端将验证固定 schema;最多四项修改", "parameterConstraints": self.storage.get_control(session["id"])["constraints"], "rejectedFeedback": rejected_feedback, "trials": [ { "number": t["number"], "rung": t["rung"], "score": t["score"], "eligible": t["eligible"], "rewardConfig": t["rewardConfig"], "evaluation": t["evaluation"] and t["evaluation"].get("metrics"), } for t in trials ], } def _fallback_patch(self, previous: dict, index: int) -> dict: names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip") name = names[index % len(names)] old = previous["weights"][name] factor = 1.1 if index % 2 == 0 else 0.9 return validate_proposal({"weights": {name: old * factor}}, previous) def _request_proposal( self, session: dict, previous: dict, base_trial_id: str, index: int ) -> dict: try: result = self.advisor.propose(self._proposal_context(session), previous) source = "agent" except Exception as error: if not session["fallbackEnabled"]: raise AdvisorUnavailable(str(error)) from error result = { "patch": self._fallback_patch(previous, index), "rationale": f"Agent 不可用,显式 fallback:{error}", "expectedImpact": {}, "confidence": 0.2, "promptHash": None, "usage": {}, "model": "optuna-fallback", } source = "fallback" 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, result["patch"], result["rationale"], result.get("expectedImpact", {}), result["confidence"], source, ) self.storage.audit( session["id"], "proposal_created", { "proposalId": proposal["id"], "source": source, "promptHash": result.get("promptHash"), "usage": result.get("usage", {}), "model": result.get("model"), }, ) return proposal def _wait_for_approval( self, session_id: str, proposal_id: str, cancel: threading.Event ) -> dict: with self.condition: while not cancel.is_set(): proposal = self.storage.get_proposal(proposal_id) if proposal["state"] != "pending": return proposal self.condition.wait(timeout=1.0) raise TuningError("session 已取消") def _run_session(self, session_id: str, resume: bool, cancel: threading.Event) -> None: try: session = self.storage.get_session(session_id) trials = self.storage.list_trials(session_id) if resume: root = self._session_root(session_id) for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]: run_dir = (root / interrupted["runDir"]).resolve() if run_dir.is_relative_to(root): shutil.rmtree(run_dir, ignore_errors=True) self.storage.delete_trial(interrupted["id"]) self.storage.audit( session_id, "session_resumed", { "discardedInterruptedTrials": [ t["id"] for t in trials if t["state"] == "interrupted" ] }, ) trials = self.storage.list_trials(session_id) if not trials: baseline_dir = "trial-000-rung-0" trial = self.storage.create_trial( session_id, 0, 0, session["config"]["rungs"][0], deepcopy(BASE_REWARD_CONFIGURATION), None, baseline_dir, ) self._execute_trial(session, trial) session = self.storage.get_session(session_id) completed_rung0 = [ t for t in self.storage.list_trials(session_id) if t["rung"] == 0 and t["state"] == "completed" ] next_number = len({t["number"] for t in completed_rung0}) best_score = max( (t["score"] 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"] and no_improve < session["config"]["earlyStopPatience"] ): if cancel.is_set(): raise TuningError("session 已取消") 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 ) # 模式允许在 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 建议" ) proposal = self._wait_for_approval(session_id, proposal["id"], cancel) if proposal["state"] == "rejected": self.storage.audit( session_id, "proposal_rejected", {"proposalId": proposal["id"], "feedback": proposal["feedback"]}, ) continue else: self.storage.decide_proposal(proposal["id"], "approved", "自动模式") proposal = self.storage.get_proposal(proposal["id"]) self._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, 0, session["config"]["rungs"][0], reward_config, proposal["id"], f"trial-{next_number:03d}-rung-0", ) result = self._execute_trial(session, trial) 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 self.storage.update_session(session_id, consecutive_no_improve=no_improve) next_number += 1 # Promote top configurations; each new rung resumes its own previous checkpoint. for rung in (1, 2): 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 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) if trial["rung"] == rung and trial["state"] == "completed" } for parent in previous[: session["config"]["promote"][rung]]: if parent["number"] in promoted_numbers: continue if cancel.is_set(): raise TuningError("session 已取消") self._claim_dispatch(session_id, cancel) root = self._session_root(session_id) checkpoint = root / parent["checkpointPath"] trial = self.storage.create_trial( session_id, parent["number"], rung, session["config"]["rungs"][rung], parent["rewardConfig"], parent["proposalId"], f"trial-{parent['number']:03d}-rung-{rung}", ) self._execute_trial(session, trial, checkpoint) self._pause_after_step(session_id) best = self._best_highest_rung(session_id) if best is None: raise TuningError("没有通过安全门槛的 trial") preset_name = f"{session['config']['runName']}-{session_id[:8]}" self.storage.save_preset(preset_name, session_id, best["id"], best["rewardConfig"]) self.storage.update_session( session_id, state="succeeded", best_trial_id=best["id"], current_trial_id=None, message="调参完成", ) self.storage.audit( session_id, "session_completed", {"bestTrialId": best["id"], "preset": preset_name} ) except Exception as error: state = self.storage.get_session(session_id)["state"] if cancel.is_set() or state == "cancelled": self.storage.update_session( session_id, state="cancelled", message="调参已取消", current_trial_id=None ) else: self.storage.update_session( session_id, state="failed", message=str(error), current_trial_id=None ) self.storage.audit(session_id, "session_failed", {"error": str(error)[:1000]}) finally: with self.lock: self.workers.pop(session_id, None) self.processes.pop(session_id, None) def _wait_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 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) if proposal["sessionId"] != session_id: 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: 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( session_id, "proposal_approved", {"proposalId": proposal_id, "modified": patch != proposal["patch"]}, ) with self.condition: self.condition.notify_all() return self.detail(session_id) def reject(self, session_id: str, proposal_id: str, payload: Any) -> dict: feedback = payload.get("feedback", "") if isinstance(payload, dict) else "" if not isinstance(feedback, str) or len(feedback) > 2000: raise TuningError("feedback 无效") proposal = self.storage.get_proposal(proposal_id) if proposal["sessionId"] != session_id: raise TuningError("proposal 不属于该 session") if not self.storage.decide_proposal(proposal_id, "rejected", feedback): raise TuningError("proposal 已处理") with self.condition: self.condition.notify_all() return self.detail(session_id) def 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="已暂停后续调度;在途 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.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: raise TuningError("当前状态不能恢复") return self.detail(session_id) def cancel(self, session_id: str) -> dict: 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: event.set() process = self.processes.get(session_id) if process: terminate_process(process) with self.condition: self.condition.notify_all() return self.detail(session_id) def metrics( self, session_id: str, trial_id: str, tags: list[str] | None, max_points: int, after_step: int | None = None, ) -> dict: trial = self.storage.get_trial(trial_id) if trial["sessionId"] != session_id: raise TuningError("trial 不属于该 session") 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) if not session["bestTrialId"]: raise TuningError("尚无最佳策略") trial = self.storage.get_trial(session["bestTrialId"]) if not trial["policyPath"]: raise TuningError("最佳策略文件不存在") root = self._session_root(session_id) path = (root / trial["policyPath"]).resolve() if not path.is_relative_to(root) or not path.is_file(): raise TuningError("最佳策略文件不存在") return path def preset_config(self, preset_id: str) -> dict: return self.storage.get_preset(preset_id)["rewardConfig"] def test_agent(self) -> dict: try: return self.advisor.test_connection() except Exception as error: raise TuningError(f"Agent 连接测试失败:{error}") from error def shutdown(self) -> None: for session_id in list(self.workers): with suppress(KeyError, TuningError): self.cancel(session_id) for worker in list(self.workers.values()): worker.join(timeout=7)