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": {}} 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_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.update_session(session["id"], state="interrupted") self.manager.resume(session["id"]) 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") if __name__ == "__main__": unittest.main()