200 lines
7.5 KiB
Python
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()
|