Download code/tt_diffusion_planner/tests/test_api_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 3.72 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_api_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_api_host.py
-
curl -L -o test_api_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_api_host.py
3.72 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """The Python API's host hooks without a device (no weights needed): ``DiffusionPlanner._prepare`` / ``_postprocess`` | |
| wired to the host code, ``__call__``'s input handling and parameter validation, ``info``, and the build refusing | |
| unknown compile parameters. | |
| TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_api_host.py | |
| """ | |
| from __future__ import annotations | |
| import threading | |
| import numpy as np | |
| import pytest | |
| from tt_diffusion_planner import io as tio | |
| from tt_diffusion_planner.api import DiffusionPlanner, Output | |
| from tt_diffusion_planner.reference import config as C | |
| from tt_diffusion_planner.tests.stubs import StubModel, sample_inputs, v5_normalization | |
| class _NoDevice(DiffusionPlanner): | |
| """The real hooks with the network replaced by the stub's fake network (no device, no weights).""" | |
| def _forward(self, prepared): | |
| return StubModel.fake_network(prepared) | |
| def _model() -> DiffusionPlanner: | |
| m = _NoDevice.__new__(_NoDevice) | |
| m._lock, m._closed, m._owns_device, m.device = threading.RLock(), False, False, None | |
| m.variant, m.compile_params, m.warm_variants, m.warmup_ms, m.device_info = "default", {}, [], {}, {} | |
| m.normalization = v5_normalization() | |
| return m | |
| def test_identity(): | |
| cls = DiffusionPlanner | |
| assert (cls.MODEL_NAME, cls.ENV_PREFIX) == ("diffusion-planner-p150", "DIFFUSION_PLANNER") | |
| assert cls.INPUT_KIND == "planner" | |
| assert cls.DEFAULT_REVISION == "423efde67f5414734da43a7ad856c17ceb8b51aa" and cls.DEFAULT_TAG == "v5.0" | |
| assert set(cls.ALLOW_PATTERNS) == {C.ENCODER_ONNX, C.DECODER_ONNX, C.TURN_INDICATOR_ONNX, C.PARAM_JSON} | |
| assert Output.__name__ == "Trajectory" and cls.LABELS == ("NONE", "DISABLE", "ENABLE_LEFT", "ENABLE_RIGHT", "KEEP") | |
| with pytest.raises(TypeError): | |
| DiffusionPlanner() | |
| def test_call_runs_the_host_hooks(): | |
| m = _model() | |
| out = m(inputs=sample_inputs(), stopping_threshold=0.5) | |
| assert isinstance(out, Output) and out.poses.shape == (80, 7) and out.predicted_agents.shape == (3, 80, 5) | |
| assert set(out.timing_ms) >= {"preprocess", "device", "postprocess", "total"} | |
| ref = StubModel()(inputs=sample_inputs(), stopping_threshold=0.5) | |
| np.testing.assert_array_equal(out.poses, ref.poses) # API hooks == the stub's direct host calls | |
| info = m.info | |
| assert info["input_kind"] == "planner" and info["input_schema"]["lanes"]["shape"] == [1, 140, 20, 33] | |
| assert info["dpm_solver_steps"] == 10 and info["input_names"] == list(C.INPUT_NAMES) | |
| def test_call_refuses_bad_requests(kwargs, match): | |
| with pytest.raises(tio.InputError, match=match): | |
| _model()(**kwargs) | |
| def test_build_refuses_unknown_compile_params(): | |
| """``from_pretrained(**compile_params)`` takes only ``precision``: any other name fails before the weights are read | |
| or the graph is built (the numerics options are ``DIFFUSION_PLANNER_*`` knobs, pinned in ``serve.env``).""" | |
| m = DiffusionPlanner.__new__(DiffusionPlanner) | |
| m.weights_path, m.compile_params, m.device = None, {"ln_fp32": "dec.*"}, None | |
| with pytest.raises(TypeError, match=r"unknown compile parameter\(s\) \['ln_fp32'\]"): | |
| DiffusionPlanner._build(m) | |