Files
uav-catch/tests/test_eval.py
robinson e901fa9579 init
2026-07-03 04:19:44 +08:00

73 lines
2.1 KiB
Python

"""Tests for control evaluation metrics."""
from __future__ import annotations
import numpy as np
from uavcatch.eval.metrics import compute_tracking_metrics, evaluate_maneuver, _step_response_metrics
from uavcatch.control.pid import circle_trajectory_with_takeoff
from uavcatch.eval.maneuvers import circle_track, step_position
def test_tracking_rmse() -> None:
err = np.array([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 0.0]])
m = compute_tracking_metrics(err)
assert abs(m.rmse[0] - np.sqrt(1 / 3)) < 1e-9
assert m.rmse_3d > 0.0
def test_step_rise_and_settle() -> None:
t = np.linspace(0, 10, 1000)
y = np.where(t < 1.0, 0.0, 1.0 - np.exp(-(t - 1.0) / 0.5))
m = _step_response_metrics(t, y, step_time=1.0, y0=0.0, yf=1.0)
assert m.rise_time is not None
assert m.settling_time is not None
assert m.overshoot_pct < 5.0
def test_evaluate_maneuver_step_spec() -> None:
spec = step_position(
np.array([0.0, 0.0, 2.0]),
np.array([0.0, 0.0, 3.0]),
step_time=2.0,
duration=10.0,
axis=2,
)
n = 200
time = np.linspace(0, 10, n)
z = np.where(time < 2.0, 2.0, 2.0 + 0.9 * (1 - np.exp(-(time - 2.0) / 1.0)))
data = {
"time": time,
"position": np.column_stack([np.zeros(n), np.zeros(n), z]),
"tracking_error": np.column_stack([np.zeros(n), np.zeros(n), 3.0 - z]),
"euler": np.zeros((n, 3)),
"omega": np.full((n, 4), 550.0),
}
metrics = evaluate_maneuver(spec, data, omega_max=838.0, mode="linear")
assert metrics.step is not None
assert metrics.tracking.rmse_3d < 0.6
def test_circle_takeoff_phase() -> None:
spec = circle_track(takeoff_time=2.0, ground_z=0.0, height=2.0, radius=2.5)
assert spec.initial_position == (0.0, 0.0, 0.05)
t0 = spec.target_fn(0.0)
assert t0.position[2] == 0.05
t1 = spec.target_fn(1.0)
assert 0.05 < t1.position[2] < 2.0
t2 = spec.target_fn(2.0)
np.testing.assert_allclose(t2.position, [2.5, 0.0, 2.0], atol=1e-9)
pos_fn, _, _ = circle_trajectory_with_takeoff(
center=np.zeros(3),
radius=2.5,
height=2.0,
period=10.0,
takeoff_time=2.0,
)
t3 = pos_fn(3.0)
assert t3.position[0] > 2.0