Files
Mujoco_WASM/training_server/tests/test_tuning_manager.py
T
chenlin deead17a9a
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
feat(training): release V0.8 自调参 Agent
2026-09-02 13:49:34 +08:00

196 lines
7.1 KiB
Python

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()