changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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)
@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)