351 lines
14 KiB
Python
351 lines
14 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": {}}
|
|
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()
|