diffusion-planner-p150 / code /tt_diffusion_planner /tests /test_host_postprocess.py
changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
10.6 kB
# SPDX-License-Identifier: Apache-2.0
"""Host post-processing (no device, no weights): denormalisation, Eigen's quaternion of the unnormalised rotation,
the trajectory velocity / force-stop / acceleration rules, predicted paths, the turn-indicator decision and the
``make_output`` assembly, against the verified research port (``dp_common.py``), hand-made cases of
``postprocessing_utils.cpp`` / ``turn_indicator_manager.cpp``, and the small goldens of the shipped samples.
TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_host_postprocess.py
"""
from __future__ import annotations
import math
import numpy as np
import pytest
from tt_diffusion_planner.host import pipeline as hp
from tt_diffusion_planner.host import postprocess as P
from tt_diffusion_planner.reference import config as C
from tt_diffusion_planner.reference.weights import load_param_json
from tt_diffusion_planner.tests import _research as R
PARAM_JSON = R.RESEARCH.parents[1] / "assets" / "diffusion-planner" / "hf_diffusion_planner" / C.PARAM_JSON
DEFAULTS = {k: spec[3] for k, spec in hp.RUNTIME_PARAMS.items()}
@pytest.fixture(scope="module")
def normalization():
if not PARAM_JSON.is_file():
pytest.skip("diffusion_planner.param.json not found")
return load_param_json(PARAM_JSON)
def _straight(n=80, v=8.0, decel_after=None, dt=0.1):
"""A straight ego path at speed v (optionally braking to a stop after ``decel_after`` points)."""
xs, x, speed = [], 0.0, v
for i in range(n):
if decel_after is not None and i >= decel_after:
speed = max(0.0, speed - 1.5)
x += speed * dt
xs.append(x)
p = np.zeros((n, 4), np.float32)
p[:, 0], p[:, 2] = xs, 1.0
return p
# ------------------------------------------------------------------------------------------------- denorm
def test_denormalize_matches_research_port(normalization):
x = np.random.default_rng(0).normal(size=(321, 81, 4)).astype(np.float32)
mean, std = normalization.state()
out = P.denormalize(x, mean, std)
np.testing.assert_array_equal(out, (x[:, 1:] * std + mean).astype(np.float32))
np.testing.assert_array_equal(out[5, 7], x[5, 8] * np.float32([20, 20, 1, 1]) + np.float32([10, 0, 0, 0]))
assert P.denormalize(x, mean, std, keep_current_state=True).shape == (321, 81, 4)
dp = R.load_script("dp_common")
if dp is not None:
sm, ss = np.asarray(normalization.state_mean), np.asarray(normalization.state_std)
np.testing.assert_array_equal(out, dp.denormalize(x[None], (sm, ss))[0])
# --------------------------------------------------------------------------------------------- orientation
def test_quaternion_of_unit_rotations_is_the_half_angle():
th = np.linspace(-math.pi + 1e-3, math.pi - 1e-3, 721)
q = P.quaternion_from_cos_sin(np.cos(th), np.sin(th))
np.testing.assert_allclose(np.abs(q[:, 3]), np.abs(np.cos(th / 2)), atol=1e-12)
np.testing.assert_allclose(np.abs(q[:, 2]), np.abs(np.sin(th / 2)), atol=1e-12)
np.testing.assert_allclose(np.linalg.norm(q, axis=1), 1.0, atol=1e-12)
np.testing.assert_allclose(P.tf2_yaw(q), th, atol=1e-12)
assert not q[:, :2].any()
def test_quaternion_of_unnormalised_rotations_follows_eigen():
"""Eigen does not normalise: |(c, s)| = r != 1 shifts the heading tf2::getYaw reads (SPEC 4.6.6). Its two
branches (trace = 2c + 1 > 0, else the m22 pivot) agree at the switch c = -0.5 only for unit vectors; for
r != 1 the published heading jumps there, a quirk of the node that is reproduced, not fixed."""
r, th = 0.9, math.pi / 2
q = P.quaternion_from_cos_sin(r * math.cos(th), r * math.sin(th))
assert math.isclose(P.tf2_yaw(q), 2 * math.atan(0.9 / (1 + 0.9 * math.cos(th))), rel_tol=1e-12)
assert abs(P.tf2_yaw(q) - th) > 0.05 # not atan2(sin, cos)
eps = 1e-12
for s in (math.sqrt(0.75), -math.sqrt(0.75)): # unit vectors: continuous
a, b = P.quaternion_from_cos_sin(-0.5 + eps, s), P.quaternion_from_cos_sin(-0.5 - eps, s)
np.testing.assert_allclose(P.tf2_yaw(a), P.tf2_yaw(b), atol=1e-9)
# r = 0.58 (c = -0.5, s = 0.3): trace branch w = 0.5, z = 0.3 -> atan2(0.3, 0.16); pivot branch
# z = sqrt(3) / 2, w = 0.3 / sqrt(3) -> atan2(0.3, -0.72)
a, b = P.quaternion_from_cos_sin(-0.5 + eps, 0.3), P.quaternion_from_cos_sin(-0.5 - eps, 0.3)
assert math.isclose(P.tf2_yaw(a), math.atan2(0.3, 0.16), rel_tol=1e-9)
assert math.isclose(P.tf2_yaw(b), math.atan2(0.3, -0.72), rel_tol=1e-9)
# --------------------------------------------------------------------------------------------- trajectory
def test_trajectory_constant_speed():
tr = P.trajectory_from_poses(_straight(v=8.0))
np.testing.assert_allclose(tr.velocity, 8.0, rtol=0, atol=1e-4)
assert tr.velocity.dtype == np.float32 and tr.acceleration.dtype == np.float32
np.testing.assert_allclose(tr.acceleration[:-1], 0.0, atol=2e-3)
assert tr.acceleration[-1] == 0 and not tr.force_stop
np.testing.assert_allclose(tr.time_from_start, 0.1 * np.arange(1, 81))
assert tr.as_columns().shape == (80, 7)
def test_trajectory_force_stop_freezes_poses():
poses = _straight(v=6.0, decel_after=20)
tr = P.trajectory_from_poses(poses, stopping_threshold=0.3)
assert tr.force_stop
stop = int(np.flatnonzero(tr.velocity == 0)[0])
assert np.all(tr.velocity[stop:] == 0)
np.testing.assert_array_equal(tr.position[stop:], np.broadcast_to(tr.position[stop - 1], tr.position[stop:].shape))
off = P.trajectory_from_poses(poses, enable_force_stop=False)
assert not off.force_stop and np.all(off.velocity[-7:] == off.velocity[80 - 8])
with pytest.raises(ValueError, match="velocity_smoothing_window"):
P.trajectory_from_poses(poses, velocity_smoothing_window=80)
@pytest.mark.skipif(R.load_script("dp_common") is None, reason="research scripts not present")
@pytest.mark.parametrize("case", ["cruise", "brake", "golden"])
def test_trajectory_matches_research_port(case):
"""``dp_common.trajectory_from_poses`` (float64 velocities, no float rounding before the average): positions after
freezing identical, velocities within float32 rounding."""
dp = R.load_script("dp_common")
if case == "golden":
p = R.FULL_GOLDENS / "straight_road.npz"
if not p.is_file():
pytest.skip("full goldens not generated")
with np.load(p) as z:
poses = z["out.trajectory"][:, [0, 1, 3, 4]]
else:
poses = _straight(v=7.0, decel_after=15 if case == "brake" else None)
tr = P.trajectory_from_poses(poses)
xy, vel, acc = dp.trajectory_from_poses(poses[:, :2])
np.testing.assert_allclose(tr.position[:, :2], xy, rtol=0, atol=1e-12)
np.testing.assert_allclose(tr.velocity, vel, rtol=1e-6, atol=1e-5)
np.testing.assert_allclose(tr.acceleration, acc, rtol=1e-5, atol=1e-3)
# ------------------------------------------------------------------------------------------ turn indicator
@pytest.mark.skipif(R.load_script("dp_common") is None, reason="research scripts not present")
def test_turn_decision_matches_research_port():
dp = R.load_script("dp_common")
rng = np.random.default_rng(4)
for _ in range(200):
lg = (rng.normal(size=5) * 4).astype(np.float32)
prev = int(rng.integers(1, 4))
d = P.TurnIndicatorManager().evaluate(lg, 0.0, prev)
cmd, prob = dp.turn_indicator_command(lg, prev_report=prev)
assert d.command == cmd
np.testing.assert_allclose(d.probabilities, prob, rtol=1e-6, atol=1e-7)
def test_turn_decision_rules_and_hold():
m = P.TurnIndicatorManager()
keep = np.array([0, 0, 0, 0, 5.0], np.float32) # KEEP wins even after -1.25
d = m.evaluate(keep, 1.0, prev_report=3)
assert (d.command, d.keep_selected) == (3, True) # repeats the last report
left = np.array([0, 0, 4.0, 0, 3.0], np.float32) # 4 > 3 - 1.25: ENABLE_LEFT
d = m.evaluate(left, 2.0, prev_report=1)
assert (d.command, d.keep_selected, d.held) == (2, False, False)
assert math.isclose(sum(d.probabilities), 1.0, rel_tol=1e-3)
d = m.evaluate(keep, 2.9, prev_report=1) # within the 1 s hold window
assert (d.command, d.held) == (2, True)
d = m.evaluate(keep, 3.01, prev_report=1) # expired
assert (d.command, d.held) == (1, False)
tie = np.array([1.0, 1.0, 0, 0, 0], np.float32) # std::max_element: the first maximum
assert P.TurnIndicatorManager().evaluate(tie).command == 0
assert P.TurnIndicatorManager().evaluate(np.zeros(0, np.float32)).command == 1
# ----------------------------------------------------------------------------------------- make_output
@pytest.mark.parametrize("stem", ["kashiwanoha_dense", "straight_road"])
def test_make_output_reproduces_small_goldens(normalization, stem):
"""From the stored reference final x0 / logits, the post-processing gives exactly the stored outputs."""
path = R.SMALL_GOLDENS / f"{stem}.outputs.npz"
if not path.is_file():
pytest.skip("small goldens not generated")
with np.load(path) as z:
g = {k: z[k] for k in z.files}
prep = hp.prepare(R.sample_raw(stem), normalization.observation)
x0 = np.zeros((C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM), np.float32)
x0[g["final_x0.rows"]] = g["final_x0"]
out = hp.make_output(x0, g["logit"], prep, normalization, DEFAULTS, model="diffusion-planner-p150")
np.testing.assert_array_equal(out.poses, g["trajectory"])
assert out.turn_indicator["command"] == int(g["turn_command"])
assert out.columns == hp.TRAJECTORY_COLUMNS and out.predicted_agents.shape == (len(prep.neighbor_rows), 80, 5)
body = out.to_dict()
assert body["num_poses"] == 80 and body["frame_id"] == "base_link"
assert body["predicted_agents"]["shape"] == [len(prep.neighbor_rows), 80, 5]
assert body["meta"]["predicted_agent_rows"] == prep.neighbor_rows.tolist()
# params: a shorter window and no force stop change only the velocity columns
alt = hp.make_output(x0, g["logit"], prep, normalization, {**DEFAULTS, "velocity_smoothing_window": 1})
np.testing.assert_array_equal(alt.poses[:, :5], out.poses[:, :5])
steps = [x0] * 11
dbg = hp.make_output(x0, g["logit"], prep, normalization, {**DEFAULTS, "return_denoising_steps": True},
denoising_steps=steps).to_dict()
assert dbg["meta"]["denoising_steps"]["shape"] == [11, 81, 4]