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