Download code/tt_diffusion_planner/tests/test_io_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 4.73 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_io_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_io_host.py
-
curl -L -o test_io_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_io_host.py
4.73 kB
| # 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 == () | |
| 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]) | |
| 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) | |
| 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) | |