changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
5.37 kB
# SPDX-License-Identifier: Apache-2.0
"""Host tests of the HTTP contract (no device, no weights): the real app of this bundle with a stub model that runs
the real host pre- and post-processing around a fake network (``tests/stubs.py``).
TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_server_host.py
"""
from __future__ import annotations
import base64
import io
import numpy as np
import pytest
fastapi = pytest.importorskip("fastapi")
from fastapi.testclient import TestClient # noqa: E402
from tt_diffusion_planner.api import DiffusionPlanner # noqa: E402
from tt_diffusion_planner.reference import config as C # noqa: E402
from tt_diffusion_planner.server import app as server # noqa: E402
from tt_diffusion_planner.tests.stubs import StubModel, sample_inputs # noqa: E402
@pytest.fixture
def client(monkeypatch):
monkeypatch.setenv("TT_MESH_SHAPE", "1x1")
monkeypatch.setenv("TT_MODEL_WEIGHTS_REVISION", "0" * 40)
monkeypatch.setattr(server.app.state.ttaw, "model_factory", StubModel)
with TestClient(server.app) as c:
yield c
def _npz_b64(arrays) -> str:
buf = io.BytesIO()
np.savez_compressed(buf, **arrays)
return base64.b64encode(buf.getvalue()).decode()
def _req(arrays=None, **extra):
return {"inputs": {"format": "npz", "data": _npz_b64(sample_inputs() if arrays is None else arrays)}, **extra}
def test_health_info_models(client):
assert client.get("/health").json()["status"] == "ok"
assert client.get("/v1/health").json()["status"] == "ok"
info = client.get("/info").json()
assert info["model"] == DiffusionPlanner.MODEL_NAME and info["weights"]["revision"] == "0" * 40
assert info["device"]["dispatch"] == "eth" and info["device"]["grid"] == "12x10"
assert info["autoware"]["package"] == "autoware_diffusion_planner"
assert info["output"]["labels"] == list(DiffusionPlanner.LABELS)
assert info["variant"] == DiffusionPlanner.DEFAULT_VARIANT
assert client.get("/v1/models").json()["data"][0]["id"] == "AutowareFoundation/diffusion_planner"
def test_predict_ok_and_equals_api(client):
r = client.post("/predict", json=_req(params={"stopping_threshold": 0.4}))
assert r.status_code == 200, r.text
body = r.json()
assert body["model"] == DiffusionPlanner.MODEL_NAME and body["frame_id"] == "base_link"
assert body["num_poses"] == 80 and len(body["trajectory"]) == 80 and len(body["trajectory"][0]) == 7
assert body["columns"] == ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"]
assert body["turn_indicator"]["command"] in (0, 1, 2, 3) and len(body["turn_indicator"]["logits"]) == 5
assert body["predicted_agents"]["shape"] == [3, 80, 5]
assert {"decode", "model_call", "total", "device"} <= set(body["timing_ms"])
state = server.app.state.ttaw
req = server.PredictRequest(**_req())
assert state.predict(req)["trajectory"] == state.model(inputs=sample_inputs()).to_dict()["trajectory"]
def test_predict_json_arrays_and_npz_output(client):
arrays = {k: v.tolist() for k, v in sample_inputs().items()}
r = client.post("/predict", json={"inputs": {"format": "json", "arrays": arrays}, "output_format": "npz"})
assert r.status_code == 200, r.text
assert r.json()["arrays"]["poses"]["shape"] == [80, 7]
def _bad(key, value):
raw = sample_inputs()
raw[key] = value
return raw
@pytest.mark.parametrize("payload,code", [
({"inputs": {"format": "npz", "data": "not base64!"}}, 400), # undecodable
(_req(_bad("lanes", np.zeros((1, 70, 20, 33), np.float32))), 400), # wrong shape
(_req(_bad("goal_pose", np.full((1, 4), np.inf, np.float32))), 400), # non-finite
(_req({k: v for k, v in sample_inputs().items() if k != "delay"}), 400), # missing tensor
(_req(params={"velocity_smoothing_window": 0}), 400), # out of range
(_req(params={"velocity_smoothing_window": 80}), 400),
(_req(params={"return_denoising_steps": "yes"}), 400), # not a boolean
(_req(params={"score_threshold": 0.5}), 400), # unknown knob
({"points": {"format": "list", "values": [[0, 0, 0, 0]]}}, 400), # wrong input kind
({}, 400), # no inputs
({"inputz": {}}, 422), # schema violation
])
def test_predict_errors(client, payload, code):
r = client.post("/predict", json=payload)
assert r.status_code == code, r.text
def test_503_while_starting(client, monkeypatch):
monkeypatch.setattr(server.app.state.ttaw, "ready", False)
assert client.post("/predict", json={}).status_code == 503
assert client.get("/health").json()["status"] == "starting"
def test_input_schema_is_published():
assert {k: list(v[0]) for k, v in DiffusionPlanner.INPUT_SCHEMA.items()} == {
k: list(s) for k, s in C.INPUT_SHAPES.items()}
def test_mesh_shape_parser():
p = server.parse_mesh_shape
assert p("1x1") == p("(1, 1)") == p("1,1") == p(None) == (1, 1)
with pytest.raises(RuntimeError):
p("two by two")