Download code/tt_diffusion_planner/api.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 7.66 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/api.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/api.py
-
curl -L -o api.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/api.py
7.66 kB
| # 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 | |