import sys import tempfile import time import unittest from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from tuning.manager import TuningManager # noqa: E402 from tuning.process import GpuLease # noqa: E402 from tuning.schema import BASE_REWARD_CONFIGURATION # noqa: E402 from tuning.scoring import score_evaluation # noqa: E402 from tuning.storage import now_iso # noqa: E402 BASE_METRICS = { "linear_velocity_rmse": 0.3, "angular_velocity_rmse": 0.2, "mean_action_acc": 0.1, "orientation_error": 0.2, "fall_rate": 0.1, "slip_velocity": 0.2, "mechanical_power": 100.0, } class FakeAdvisor: def capability(self): return { "configured": True, "apiKeyConfigured": True, "frameworkInstalled": True, "model": "fake", "baseUrl": "https://example.invalid", } def propose(self, _context, previous): value = min(2.4, previous["weights"]["pose"] * 1.05) return { "patch": {"weights": {"pose": value}, "params": {}}, "rationale": "fake", "expectedImpact": {}, "confidence": 0.8, "promptHash": "abc", "usage": {}, "model": "fake", } def test_connection(self): return {"ok": True, "model": "fake", "outputType": "fake"} class FakeTuningManager(TuningManager): def _execute_trial(self, session, trial, resume_checkpoint=None): del resume_checkpoint factor = max(0.5, 1.0 - 0.03 * trial["number"] - 0.01 * trial["rung"]) metrics = {key: value * factor for key, value in BASE_METRICS.items()} trials = self.storage.list_trials(session["id"]) baseline = next( (item for item in trials if item["number"] == 0 and item["rung"] == 0), None ) if baseline and baseline["evaluation"]: scored = score_evaluation( baseline["evaluation"]["metrics"], metrics, session["objectiveWeights"] ) else: scored = {"score": 0.0, "eligible": True, "components": {}} root = self._session_root(session["id"]) run = root / trial["runDir"] run.mkdir(parents=True, exist_ok=True) (run / "model_1.pt").write_bytes(b"checkpoint") (run / "policy.onnx").write_bytes(b"onnx") evaluation = {"metrics": metrics, "score": scored} self.storage.update_trial( trial["id"], state="completed", started_at=now_iso(), ended_at=now_iso(), message="fake complete", checkpoint_path=str((run / "model_1.pt").relative_to(root)), policy_path=str((run / "policy.onnx").relative_to(root)), evaluation=evaluation, score=scored["score"], eligible=scored["eligible"], ) return self.storage.get_trial(trial["id"]) class TuningManagerTest(unittest.TestCase): def setUp(self): self.temporary = tempfile.TemporaryDirectory() self.root = Path(self.temporary.name) (self.root / "trainer" / "scripts").mkdir(parents=True) (self.root / "trainer" / "scripts" / "evaluate.py").write_text("", encoding="utf-8") self.manager = FakeTuningManager( self.root / "trainer", sys.executable, self.root / "data", GpuLease(), advisor=FakeAdvisor(), ) def tearDown(self): self.manager.shutdown() self.temporary.cleanup() @staticmethod def payload(mode="automatic"): return { "taskId": "Unitree-Go2-Flat", "mode": mode, "runName": "test", "numEnvs": 16, "gpuIds": [0], "trialCount": 4, "initialIterations": 1, "middleIterations": 2, "finalIterations": 3, "evalNumEnvs": 8, "evalSteps": 10, } def wait_terminal(self, session_id, timeout=5): deadline = time.monotonic() + timeout while time.monotonic() < deadline: session = self.manager.detail(session_id) if session["state"] in {"succeeded", "failed", "cancelled"}: return session time.sleep(0.01) self.fail("session did not finish") def test_automatic_session_runs_rungs_and_persists_best_artifact(self): session = self.manager.create(self.payload()) completed = self.wait_terminal(session["id"]) self.assertEqual(completed["state"], "succeeded", completed["message"]) self.assertGreaterEqual(len(completed["trials"]), 7) self.assertTrue(self.manager.best_artifact(session["id"]).is_file()) self.assertEqual(len(self.manager.storage.list_presets()), 1) def test_approval_session_waits_and_accepts_modified_patch(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] 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} ) self.assertEqual(approved["proposals"][-1]["state"], "approved") 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) baseline = self.manager.storage.create_trial( session["id"], 0, 0, 1, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0", ) self.manager._execute_trial(self.manager.storage.get_session(session["id"]), baseline) interrupted = self.manager.storage.create_trial( session["id"], 1, 0, 1, baseline["rewardConfig"], None, "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"]]) def test_create_validation_and_agent_capability(self): self.assertTrue(self.manager.capability()["configured"]) with self.assertRaisesRegex(Exception, "只支持"): 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()