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)