File size: 4,731 Bytes
4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 | # SPDX-License-Identifier: Apache-2.0
"""Host tests (no device, no weights): the request decoders shared by the API and the server, i.e. the vendored
``ttaw.io`` as this bundle binds it (``io.load_inputs`` / ``io.decode_inputs`` with ``INPUT_SCHEMA``; the full suite of
ttaw.io is common/tests/host/test_io_outputs_host.py).
TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_io_host.py
"""
from __future__ import annotations
import base64
import io
from pathlib import Path
import numpy as np
import pytest
from tt_diffusion_planner import io as tio
from tt_diffusion_planner.api import DiffusionPlanner
from tt_diffusion_planner.reference import config as C
from tt_diffusion_planner.tests.stubs import sample_inputs
SAMPLES = ("kashiwanoha_dense", "straight_road")
def _npz_b64(arrays) -> str:
buf = io.BytesIO()
np.savez_compressed(buf, **arrays)
return base64.b64encode(buf.getvalue()).decode("ascii")
def test_schema_is_the_node_input_map():
"""The 15 raw tensors of ``DiffusionPlannerCore::create_input_data`` (core.cpp:414-595), batch 1, float32."""
assert DiffusionPlanner.INPUT_SCHEMA is C.INPUT_SCHEMA and DiffusionPlanner.INPUT_KIND == "planner"
assert set(C.INPUT_SCHEMA) == {
"sampled_trajectories", "ego_agent_past", "ego_current_state", "neighbor_agents_past", "static_objects",
"lanes", "lanes_speed_limit", "route_lanes", "route_lanes_speed_limit", "polygons", "line_strings",
"goal_pose", "ego_shape", "turn_indicators", "delay"}
assert C.INPUT_SCHEMA["neighbor_agents_past"][0] == (1, 320, 31, 11)
assert C.INPUT_SCHEMA["lanes"][0] == (1, 140, 20, 33) and C.INPUT_SCHEMA["line_strings"][0] == (1, 60, 20, 4)
assert all(np.dtype(dt) == np.float32 for _, dt in C.INPUT_SCHEMA.values())
assert tio.DEFAULT_POINT_FIELDS == () and DiffusionPlanner.POINT_FIELDS == ()
@pytest.mark.parametrize("stem", SAMPLES)
def test_shipped_samples_load_in_every_form(stem):
path = Path(__file__).resolve().parents[1] / "samples" / f"{stem}.npz"
arrays = tio.load_inputs(str(path))
assert set(arrays) == set(C.INPUT_NAMES) and all(a.dtype == np.float32 for a in arrays.values())
for form in (path.read_bytes(), {k: v.astype(np.float64) for k, v in arrays.items()},
{"format": "npz", "data": base64.b64encode(path.read_bytes()).decode()}):
again = tio.load_inputs(form)
for k in C.INPUT_NAMES:
np.testing.assert_array_equal(again[k], arrays[k])
dec = tio.decode_inputs({"format": "npz", "data": _npz_b64(arrays)})
np.testing.assert_array_equal(dec["lanes"], arrays["lanes"])
def test_json_arrays_form():
raw = sample_inputs()
spec = {"format": "json", "arrays": {k: v.tolist() for k, v in raw.items()}}
got = tio.decode_inputs(spec)
for k in C.INPUT_NAMES:
np.testing.assert_array_equal(got[k], raw[k])
@pytest.mark.parametrize("mutate,match", [
(lambda r: r.pop("delay"), "missing"),
(lambda r: r.update(extra=np.zeros(3, np.float32)), "unexpected"),
(lambda r: r.update(lanes=np.zeros((1, 70, 20, 33), np.float32)), "shape"),
(lambda r: r.update(ego_shape=np.zeros(3, np.float32)), "shape"),
(lambda r: r["neighbor_agents_past"].__setitem__((0, 0, 0, 0), np.nan), "non-finite"),
(lambda r: r["goal_pose"].__setitem__((0, 0), np.inf), "non-finite"),
])
def test_invalid_inputs_raise_input_error(mutate, match):
raw = sample_inputs()
mutate(raw)
with pytest.raises(tio.InputError, match=match):
tio.load_inputs(raw)
with pytest.raises(tio.InputError, match=match):
tio.decode_inputs({"format": "npz", "data": _npz_b64(raw)})
def test_undecodable_envelopes():
with pytest.raises(tio.InputError):
tio.decode_inputs({"format": "npz", "data": "not base64!"})
with pytest.raises(tio.InputError):
tio.decode_inputs({"format": "npz", "data": base64.b64encode(b"not an npz").decode()})
with pytest.raises(tio.InputError, match="format"):
tio.decode_inputs({"format": "pickle", "data": ""})
with pytest.raises(tio.InputError):
tio.load_inputs(42)
@pytest.mark.parametrize("fmt", ["npz", "npz_compressed", "raw", "list"])
def test_encode_array_roundtrip(fmt):
a = np.arange(2 * 80 * 5, dtype=np.float32).reshape(2, 80, 5)
d = tio.encode_array(a, fmt=fmt, key="predicted_agents")
if fmt.startswith("npz"):
back = np.load(io.BytesIO(base64.b64decode(d["data"])))["predicted_agents"]
elif fmt == "raw":
back = np.frombuffer(base64.b64decode(d["data"]), dtype=d["dtype"]).reshape(d["shape"])
else:
back = np.asarray(d["data"], dtype=d["dtype"])
np.testing.assert_array_equal(back, a)
|