File size: 3,721 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
# 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)