240 lines
9.4 KiB
Python
240 lines
9.4 KiB
Python
import math
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
|
|
|
from tuning.advisor import AdvisorConfig, DeepSeekAdvisor # noqa: E402
|
|
from tuning.schema import ( # noqa: E402
|
|
BASE_REWARD_CONFIGURATION,
|
|
RewardConfigError,
|
|
merge_proposal,
|
|
validate_configuration,
|
|
validate_constraints,
|
|
validate_proposal,
|
|
)
|
|
from tuning.scoring import ( # noqa: E402
|
|
DEFAULT_OBJECTIVE_WEIGHTS,
|
|
EvaluationError,
|
|
score_evaluation,
|
|
)
|
|
from tuning.storage import StorageConflict, TuningStorage # noqa: E402
|
|
|
|
|
|
class RewardSchemaTest(unittest.TestCase):
|
|
def test_baseline_is_complete_and_energy_is_disabled(self):
|
|
config = validate_configuration(BASE_REWARD_CONFIGURATION)
|
|
self.assertEqual(len(config["weights"]), 16)
|
|
self.assertEqual(config["weights"]["electrical_power"], 0.0)
|
|
|
|
def test_sparse_proposal_constraints(self):
|
|
patch = validate_proposal(
|
|
{
|
|
"weights": {"track_linear_velocity": 1.2, "foot_slip": -0.3},
|
|
"params": {"foot_gait.period": 0.65},
|
|
},
|
|
BASE_REWARD_CONFIGURATION,
|
|
)
|
|
merged = merge_proposal(BASE_REWARD_CONFIGURATION, patch)
|
|
self.assertEqual(merged["weights"]["track_linear_velocity"], 1.2)
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal({"weights": {"track_linear_velocity": 0}}, BASE_REWARD_CONFIGURATION)
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal({"weights": {"foot_slip": 0.2}}, BASE_REWARD_CONFIGURATION)
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal(
|
|
{"weights": {"track_linear_velocity": 3.0}}, BASE_REWARD_CONFIGURATION
|
|
)
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal({"params": {"foot_gait.period": math.nan}}, BASE_REWARD_CONFIGURATION)
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal(
|
|
{
|
|
"weights": {
|
|
"pose": 1.1,
|
|
"foot_gait": 0.6,
|
|
"foot_slip": -0.3,
|
|
"soft_landing": -0.002,
|
|
},
|
|
"params": {"foot_gait.period": 0.65},
|
|
},
|
|
BASE_REWARD_CONFIGURATION,
|
|
)
|
|
|
|
def test_cross_parameter_order(self):
|
|
with self.assertRaises(RewardConfigError):
|
|
validate_proposal(
|
|
{"params": {"pose.walking_threshold": 0.5, "pose.running_threshold": 0.4}},
|
|
BASE_REWARD_CONFIGURATION,
|
|
)
|
|
|
|
def test_session_constraints_reject_unknown_out_of_range_and_fixed_changes(self):
|
|
constraints = validate_constraints(
|
|
{
|
|
"weights.track_linear_velocity": {"kind": "fixed", "value": 1.0},
|
|
"params.foot_gait.period": {"kind": "range", "min": 0.5, "max": 0.7},
|
|
}
|
|
)
|
|
validate_proposal(
|
|
{"params": {"foot_gait.period": 0.65}},
|
|
BASE_REWARD_CONFIGURATION,
|
|
constraints,
|
|
)
|
|
with self.assertRaisesRegex(RewardConfigError, "已固定"):
|
|
validate_proposal(
|
|
{"weights": {"track_linear_velocity": 1.1}},
|
|
BASE_REWARD_CONFIGURATION,
|
|
constraints,
|
|
)
|
|
with self.assertRaisesRegex(RewardConfigError, "工程锁定范围"):
|
|
validate_proposal(
|
|
{"params": {"foot_gait.period": 0.75}},
|
|
BASE_REWARD_CONFIGURATION,
|
|
constraints,
|
|
)
|
|
with self.assertRaisesRegex(RewardConfigError, "未知参数约束"):
|
|
validate_constraints({"weights.not_allowed": {"kind": "fixed", "value": 1.0}})
|
|
|
|
|
|
class ScoringTest(unittest.TestCase):
|
|
baseline = {
|
|
"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,
|
|
}
|
|
|
|
def test_improvement_and_safety_gate(self):
|
|
better = {key: value * 0.8 for key, value in self.baseline.items()}
|
|
scored = score_evaluation(self.baseline, better, DEFAULT_OBJECTIVE_WEIGHTS)
|
|
self.assertTrue(scored["eligible"])
|
|
self.assertGreater(scored["score"], 0)
|
|
unsafe = dict(better, fall_rate=0.2)
|
|
scored = score_evaluation(self.baseline, unsafe)
|
|
self.assertFalse(scored["eligible"])
|
|
self.assertEqual(scored["score"], -1.0)
|
|
|
|
def test_rejects_missing_and_nonfinite_metrics(self):
|
|
with self.assertRaises(EvaluationError):
|
|
score_evaluation(self.baseline, {"fall_rate": 0.1})
|
|
bad = dict(self.baseline, mechanical_power=math.inf)
|
|
with self.assertRaises(EvaluationError):
|
|
score_evaluation(self.baseline, bad)
|
|
|
|
|
|
class AdvisorTest(unittest.TestCase):
|
|
class Output:
|
|
weights = {"pose": 1.1}
|
|
params = {}
|
|
rationale = "improve posture"
|
|
expected_impact = {"posture": "better"}
|
|
confidence = 0.7
|
|
|
|
class Result:
|
|
output = None
|
|
|
|
@staticmethod
|
|
def usage():
|
|
return type("Usage", (), {"requests": 1, "input_tokens": 10, "output_tokens": 5})()
|
|
|
|
class Agent:
|
|
def __init__(self, output):
|
|
self.output = output
|
|
|
|
def run_sync(self, _prompt):
|
|
result = AdvisorTest.Result()
|
|
result.output = self.output
|
|
return result
|
|
|
|
def test_structured_result_is_revalidated_locally(self):
|
|
advisor = DeepSeekAdvisor(AdvisorConfig("fake"))
|
|
advisor._cached_agent = self.Agent(self.Output())
|
|
proposal = advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION)
|
|
self.assertEqual(proposal["patch"]["weights"]["pose"], 1.1)
|
|
self.assertEqual(proposal["usage"]["input_tokens"], 10)
|
|
|
|
def test_invalid_model_patch_is_rejected(self):
|
|
output = self.Output()
|
|
output.weights = {"track_linear_velocity": -1.0}
|
|
advisor = DeepSeekAdvisor(AdvisorConfig("fake"))
|
|
advisor._cached_agent = self.Agent(output)
|
|
with self.assertRaises(RewardConfigError):
|
|
advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION)
|
|
|
|
|
|
class StorageTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temporary = tempfile.TemporaryDirectory()
|
|
self.storage = TuningStorage(Path(self.temporary.name) / "state.sqlite3")
|
|
|
|
def tearDown(self):
|
|
self.temporary.cleanup()
|
|
|
|
def test_persists_session_trial_proposal_and_metrics(self):
|
|
session = self.storage.create_session(
|
|
"approval", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, False
|
|
)
|
|
trial = self.storage.create_trial(
|
|
session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0"
|
|
)
|
|
proposal = self.storage.create_proposal(
|
|
session["id"],
|
|
trial["id"],
|
|
{"weights": {"pose": 1.1}, "params": {}},
|
|
"test",
|
|
{},
|
|
0.8,
|
|
)
|
|
self.assertTrue(self.storage.decide_proposal(proposal["id"], "approved", None))
|
|
self.assertFalse(self.storage.decide_proposal(proposal["id"], "approved", None))
|
|
points = [
|
|
("Train/reward", step, float(step), 50.0 if step == 50 else float(step % 7))
|
|
for step in range(100)
|
|
]
|
|
self.storage.insert_metrics(trial["id"], points)
|
|
sampled = self.storage.metrics(trial["id"], max_points=10)[0]["points"]
|
|
self.assertEqual(len(sampled), 10)
|
|
self.assertEqual(sampled[0]["step"], 0)
|
|
self.assertEqual(sampled[-1]["step"], 99)
|
|
self.assertIn(50.0, [point["value"] for point in sampled])
|
|
control = self.storage.replace_constraints(
|
|
session["id"],
|
|
0,
|
|
{"weights.pose": {"kind": "range", "min": 0.5, "max": 1.5}},
|
|
)
|
|
self.assertEqual(control["constraintsRevision"], 1)
|
|
with self.assertRaises(StorageConflict):
|
|
self.storage.replace_constraints(session["id"], 0, {})
|
|
self.storage.grant_dispatch_token(session["id"])
|
|
with self.assertRaises(StorageConflict):
|
|
self.storage.grant_dispatch_token(session["id"])
|
|
self.assertTrue(self.storage.use_dispatch_token(session["id"]))
|
|
self.assertFalse(self.storage.use_dispatch_token(session["id"]))
|
|
incremental = self.storage.metrics(trial["id"], max_points=100, after_step=90)[0]
|
|
self.assertEqual(incremental["points"][0]["step"], 91)
|
|
reopened = TuningStorage(self.storage.path)
|
|
self.assertEqual(reopened.get_session(session["id"])["mode"], "approval")
|
|
self.assertEqual(reopened.get_control(session["id"])["constraintsRevision"], 1)
|
|
|
|
def test_recovery_marks_inflight_records(self):
|
|
session = self.storage.create_session(
|
|
"automatic", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, True
|
|
)
|
|
trial = self.storage.create_trial(
|
|
session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0"
|
|
)
|
|
self.storage.update_session(session["id"], state="running")
|
|
self.storage.update_trial(trial["id"], state="training")
|
|
self.storage.recover_interrupted()
|
|
self.assertEqual(self.storage.get_session(session["id"])["state"], "interrupted")
|
|
self.assertEqual(self.storage.get_trial(trial["id"])["state"], "interrupted")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|