# 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) @pytest.mark.parametrize("kwargs,match", [ ({}, "needs `inputs`"), ({"points": np.zeros((4, 4), np.float32), "inputs": sample_inputs()}, "only `inputs`"), ({"inputs": {k: v for k, v in sample_inputs().items() if k != "lanes"}}, "missing"), ({"inputs": sample_inputs(), "velocity_smoothing_window": 0}, "outside"), ({"inputs": sample_inputs(), "return_denoising_steps": 1}, "boolean"), ({"inputs": sample_inputs(), "score_threshold": 0.3}, "unknown parameter"), ]) 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)