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