# 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