exoskeleton/code/test/test_retargeting_baselines.py

200 lines
7.5 KiB
Python

"""Contract tests for the three formal H1 retargeting baselines."""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
import numpy as np
CODE_ROOT = Path(__file__).resolve().parents[1]
if str(CODE_ROOT) not in sys.path:
sys.path.insert(0, str(CODE_ROOT))
from core.model_contract import ( # noqa: E402
MASTER_FRAMES,
MASTER_JOINT_NAMES,
SLAVE_FRAMES,
SLAVE_JOINT_NAMES,
joint_q_indices,
load_models,
)
from core.retargeting_baselines import ( # noqa: E402
Retargeter,
RetargetingFailure,
RetargetingSolverStatus,
SEWRetargeterAdapter,
build_canonical_baselines,
build_canonical_sew_target_baselines,
)
from core.sew_mapper2 import BallJointConfig, SEWMapper # noqa: E402
class RetargetingBaselinesTest(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.models = load_models()
cls.baselines = build_canonical_baselines(cls.models)
cls.slave_indices = joint_q_indices(
cls.models.slave, SLAVE_JOINT_NAMES
)
def test_factory_exposes_three_uniform_methods(self) -> None:
self.assertEqual(
set(self.baselines),
{"scaled_joint_space", "bounded_dls_ik", "task_priority_ik"},
)
for baseline in self.baselines.values():
self.assertIsInstance(baseline, Retargeter)
def test_frozen_reference_is_exact_for_all_methods(self) -> None:
q_master = np.zeros(7)
for name, baseline in self.baselines.items():
with self.subTest(method=name):
result = baseline.retarget(q_master)
self.assertTrue(result.success, result.events)
self.assertTrue(result.smooth, result.events)
self.assertIs(result.failure, RetargetingFailure.NONE)
self.assertEqual(result.q_slave.shape, (self.models.slave.nq,))
self.assertFalse(result.q_slave.flags.writeable)
self.assertIsNotNone(result.target)
self.assertGreaterEqual(result.diagnostics.runtime_s, 0.0)
self.assertGreaterEqual(result.diagnostics.cost, 0.0)
def test_cartesian_baselines_share_target_and_converge(self) -> None:
q_master = np.array(
[0.05, -0.04, 0.03, 0.08, -0.02, 0.03, -0.01],
dtype=float,
)
dls = self.baselines["bounded_dls_ik"].retarget(q_master)
priority = self.baselines["task_priority_ik"].retarget(q_master)
self.assertTrue(dls.success, dls.events)
self.assertTrue(priority.success, priority.events)
self.assertIs(
dls.diagnostics.status, RetargetingSolverStatus.CONVERGED
)
self.assertIs(
priority.diagnostics.status, RetargetingSolverStatus.CONVERGED
)
np.testing.assert_allclose(
dls.target.position, priority.target.position, atol=0.0
)
np.testing.assert_allclose(
dls.target.rotation, priority.target.rotation, atol=0.0
)
self.assertLess(dls.diagnostics.position_error_m, 2e-4)
self.assertLess(dls.diagnostics.orientation_error_rad, 2e-3)
self.assertLess(priority.diagnostics.position_error_m, 2e-4)
self.assertLess(priority.diagnostics.orientation_error_rad, 2e-3)
self.assertEqual(
priority.diagnostics.message,
"secondary_objective=log_manipulability",
)
def test_scaled_joint_mapping_uses_declared_urdf_ranges(self) -> None:
baseline = self.baselines["scaled_joint_space"]
result = baseline.retarget(np.zeros(7))
expected_midpoint = 0.5 * (
baseline.slave_lower + baseline.slave_upper
)
np.testing.assert_allclose(
result.q_slave[self.slave_indices], expected_midpoint, atol=1e-15
)
self.assertIs(
result.diagnostics.status, RetargetingSolverStatus.CLOSED_FORM
)
def test_master_limit_failure_is_returned_not_thrown(self) -> None:
q_outside = np.zeros(7)
q_outside[0] = np.pi + 0.01
for name, baseline in self.baselines.items():
with self.subTest(method=name):
result = baseline.retarget(q_outside)
self.assertFalse(result.success)
self.assertFalse(result.smooth)
self.assertIs(
result.failure,
RetargetingFailure.MASTER_LIMIT_VIOLATION,
)
self.assertIs(
result.diagnostics.status,
RetargetingSolverStatus.INVALID_INPUT,
)
self.assertIn("master_limit_violation", result.events)
def test_invalid_seed_has_stable_failure(self) -> None:
for name in ("bounded_dls_ik", "task_priority_ik"):
result = self.baselines[name].retarget(
np.zeros(7), q_slave_seed=np.zeros(3)
)
self.assertIs(result.failure, RetargetingFailure.INVALID_INPUT)
self.assertEqual(result.events, ("invalid_input",))
def test_existing_sew_mapper_has_compatible_adapter(self) -> None:
mapper = SEWMapper(
master_model=self.models.master,
slave_model=self.models.slave,
m_shoulder=MASTER_FRAMES["shoulder"],
m_elbow=MASTER_FRAMES["elbow"],
m_wrist=MASTER_FRAMES["wrist"],
m_ee=MASTER_FRAMES["ee"],
s_shoulder=SLAVE_FRAMES["shoulder"],
s_elbow=SLAVE_FRAMES["elbow"],
s_wrist=SLAVE_FRAMES["wrist"],
s_ee=SLAVE_FRAMES["wrist"],
master_joint_names=MASTER_JOINT_NAMES,
slave_joint_names=SLAVE_JOINT_NAMES,
slave_shoulder_cfg=BallJointConfig(
"yxy", SLAVE_JOINT_NAMES[:3], (-1.0, 1.0, -1.0)
),
slave_wrist_cfg=BallJointConfig(
"yzx", SLAVE_JOINT_NAMES[4:], (-1.0, 1.0, 1.0)
),
)
adapter = SEWRetargeterAdapter(mapper)
result = adapter.retarget(
np.array([0.534, 0.314, -0.1, 2.14, 0.38, 0.38, -0.72])
)
self.assertTrue(result.success, result.events)
self.assertEqual(result.method, "sew")
self.assertIs(
result.diagnostics.status, RetargetingSolverStatus.CONVERGED
)
invalid = adapter.retarget(np.zeros(3))
self.assertIs(invalid.failure, RetargetingFailure.INVALID_INPUT)
def test_common_sew_target_factory_shares_wrist_target(self) -> None:
baselines, proposed = build_canonical_sew_target_baselines(self.models)
q_master = np.array(
[0.534, 0.314, -0.1, 2.14, 0.38, 0.38, -0.72]
)
dls = baselines["bounded_dls_ik"].retarget(q_master)
priority = baselines["task_priority_ik"].retarget(q_master)
sew = proposed.retarget(q_master)
self.assertTrue(dls.success, dls.events)
self.assertTrue(priority.success, priority.events)
self.assertTrue(sew.success, sew.events)
np.testing.assert_allclose(
dls.target.position, priority.target.position, atol=0.0
)
np.testing.assert_allclose(
dls.target.rotation, priority.target.rotation, atol=0.0
)
degenerate = baselines["bounded_dls_ik"].retarget(
np.array([-np.pi / 2.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0])
)
self.assertIs(
degenerate.failure, RetargetingFailure.DEGENERATE_GEOMETRY
)
self.assertIn("master_arm_plane_degenerate", degenerate.events)
if __name__ == "__main__":
unittest.main()