File size: 7,656 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
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
# SPDX-License-Identifier: Apache-2.0
"""Python API: Diffusion Planner v5.0 (Autoware diffusion_planner) on one Tenstorrent Blackhole p150.

    from tt_diffusion_planner import DiffusionPlanner

    with DiffusionPlanner.from_pretrained(device_id=0) as model:  # weights -> HF cache, device open, traces captured
        out = model(inputs="code/tt_diffusion_planner/samples/kashiwanoha_dense.npz")
    print(out.to_dict())  # the same JSON as POST /predict

The contract shared by every bundle of the Autoware collection is the vendored ``ttaw.api_base.ModelBase``
(BUNDLE_CONVENTIONS.md section 8): ``from_pretrained`` resolves the pinned weights before it claims the chip, opens it
(ETH dispatch, 12x10), builds the graph and captures every trace variant in ``warmup_variants``, so the first call is
as fast as the later ones; calls are serialised by a lock (one chip, batch 1) and fill ``timing_ms``; ``close()`` is
idempotent, also runs at interpreter exit, and closes the chip only if the model opened it. The HTTP server
(``tt_diffusion_planner.server.app``) calls this class, so ``/predict`` and ``model(...)`` agree bit for bit.

Input: ``inputs=`` holds the 15 raw tensors of the node's ``DiffusionPlannerCore::create_input_data`` (ego frame,
before normalization, batch 1; ``INPUT_SCHEMA``). Converting ROS messages and the Lanelet2 map into them, and the
node's temporal state (agent buffers, ego history, RTC prefix of ``sampled_trajectories``), stay with the client.
The hooks delegate to ``tt_diffusion_planner.host`` (the node's pre- and post-processing, numpy) around
``tt_diffusion_planner.tt`` (the ttnn graph: encoder + 11 DiT evaluations with the DPM-Solver++(2M) update + turn head,
run by a ``ttaw.trace.TraceRunner``). Importing this module has no side effects (ttnn / torch only inside hooks).
"""
from __future__ import annotations

from typing import Any, Dict

from . import io as tio
from .device import DEVICE_DEFAULTS
from .host import pipeline as hp
from .reference import config as C
from .ttaw.api_base import ModelBase
from .ttaw.outputs import Trajectory

__all__ = ["DiffusionPlanner", "Output"]

# The result class of this model (ttaw.outputs: Detections3D, Detections2D, Segmentation3D, Mask2D, Trajectory).
Output = Trajectory


class DiffusionPlanner(ModelBase):
    """Diffusion Planner v5.0 (Autoware diffusion_planner) on one Blackhole p150. Create it with
    :meth:`from_pretrained`."""

    MODEL_NAME = "diffusion-planner-p150"
    ENV_PREFIX = "DIFFUSION_PLANNER"  # prefix of the environment knobs (SERVING.md section 3.4)
    DEFAULT_REPO = "AutowareFoundation/diffusion_planner"
    DEFAULT_TAG = "v5.0"  # the Autoware ansible artifacts pin; the node loads only major version 5
    # the commit DEFAULT_TAG points to (pinned: tags can move)
    DEFAULT_REVISION = "423efde67f5414734da43a7ad856c17ceb8b51aa"
    ALLOW_PATTERNS = ["diffusion_planner_encoder.onnx", "diffusion_planner_decoder.onnx",
                      "diffusion_planner_turn_indicator.onnx", "diffusion_planner.param.json"]
    VARIANTS = ["default"]  # load-time: the multi-step graph with dpm_solver_steps = 10
    DEFAULT_VARIANT = "default"
    INPUT_KIND = "planner"  # lidar | camera | multicam | lidar+multicam | planner
    CAMERA_ORDER = ()
    POINT_FIELDS = tio.DEFAULT_POINT_FIELDS  # () : no point cloud
    # turn-indicator logit order (dimensions.hpp:74-79); the published command is the index for 0..3
    LABELS = C.TURN_INDICATOR_LABELS
    # Per-request knobs: name -> (type, min, max, default); host-side post-processing only (the node's YAML defaults).
    RUNTIME_PARAMS = hp.RUNTIME_PARAMS
    EXTRA_INPUTS = ()
    # the ONNX-named raw tensors (name -> (shape, dtype)), decoded and checked on every call (API and server alike)
    INPUT_SCHEMA = C.INPUT_SCHEMA
    DEVICE_DEFAULTS = DEVICE_DEFAULTS  # validated open parameters (device.py)

    # ---- port-specific hooks (called by ModelBase; keep host work out of _forward) ---------------------------
    def _build(self) -> None:
        """Weights (the three ONNX files + param JSON, read as data by ``reference.weights``) -> the ttnn graph of
        ``tt_diffusion_planner.tt`` registered as the variants of a ``ttaw.trace.TraceRunner`` (persistent inputs,
        RT-dev solver / adaLN tables and states allocated here, before any capture). No capture here."""
        from .reference.weights import load_weights
        from .tt.model import TtDiffusionPlanner

        unknown = sorted(set(self.compile_params) - {"precision"})
        if unknown:
            raise TypeError(f"unknown compile parameter(s) {unknown}; allowed: precision (extra precision rules, "
                            "e.g. 'dec.*=HiFi2+fp32'); the LN_FP32 / HIDDEN_FP32 options are DIFFUSION_PLANNER_* knobs")
        self.planner_weights = load_weights(self.weights_path)
        self.normalization = self.planner_weights.normalization
        self.tt = TtDiffusionPlanner(self.device, self.planner_weights, precision=self.compile_params.get("precision"))
        self.runner = self.tt.runner

    def _warm_one(self, variant: Dict[str, Any]) -> None:
        """``TraceRunner.capture`` warms every pending variant eagerly (kernel JIT, program cache) before any capture,
        then captures with program-cache misses forbidden; idempotent."""
        self.runner.capture()

    def _prepare(self, points: Any = None, inputs: Any = None, **other: Any) -> hp.Prepared:
        """The node's host pre-processing (``host.prepare``): normalization (all-zero rows kept), speed masks, the
        encoder's host features, the decoder masks and the solver's initial state."""
        given = sorted(k for k, v in {"points": points, **other}.items() if v is not None)
        if given:
            raise tio.InputError(f"this model takes only `inputs` (the planner tensors), not {given}")
        if inputs is None:
            raise tio.InputError("this model needs `inputs`: the 15 planner tensors of INPUT_SCHEMA")
        return hp.prepare(inputs, self.normalization.observation)

    def _forward(self, prepared: hp.Prepared) -> Dict[str, Any]:
        """Upload the host features into the persistent device inputs, replay the plan's trace(s) and read back
        ``final_x0`` (normalised [321, 81, 4]), the turn logits and, when asked, the solver iterates."""
        return self.tt.forward(prepared)

    def _postprocess(self, raw: Dict[str, Any], prepared: hp.Prepared, params: Dict[str, Any]) -> Output:
        """The node's post-processing (``host.make_output``): denormalisation, trajectory velocity / force-stop /
        acceleration, predicted neighbour paths, turn-indicator decision."""
        return hp.make_output(raw["final_x0"], raw["logit"], prepared, self.normalization, params,
                              model=self.MODEL_NAME, denoising_steps=raw.get("denoising_steps"))

    def _release(self) -> None:
        """Release the traces and persistent device tensors; also called when ``from_pretrained`` fails half-way."""
        runner = getattr(self, "runner", None)
        if runner is not None:
            runner.release()

    def extra_info(self) -> Dict[str, Any]:
        """Additions to ``model.info`` and ``/info``: the trace variants, CQs and persistent tensors of the runner."""
        runner = getattr(self, "runner", None)
        info: Dict[str, Any] = {"dpm_solver_steps": C.DPM_SOLVER_STEPS, "input_names": list(C.INPUT_NAMES)}
        tt = getattr(self, "tt", None)
        if tt is not None:
            info.update(tt.describe())
        elif runner is not None:
            info["trace"] = runner.describe()
        return info