Download code/tt_diffusion_planner/ttaw/api_base.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 21.1 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/api_base.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/api_base.py
-
curl -L -o api_base.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/api_base.py
21.1 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """C08 ``ModelBase``: the Python API contract shared by every bundle (BUNDLE_CONVENTIONS.md section 8). | |
| A bundle's ``api.py`` subclasses :class:`ModelBase`, sets the identity class attributes and implements six hooks:: | |
| class CenterPoint(ModelBase): | |
| MODEL_NAME = "centerpoint-p150" | |
| ENV_PREFIX = "CENTERPOINT" | |
| DEFAULT_REPO = "AutowareFoundation/lidar_centerpoint" | |
| DEFAULT_TAG, DEFAULT_REVISION = "v4.1", "494c8171def40bd36cc2feb323e0a5acbfab132b" | |
| ALLOW_PATTERNS = ["base/*", "tiny/*"] | |
| VARIANTS, DEFAULT_VARIANT = ("base", "tiny"), "base" | |
| LABELS = ("CAR", "TRUCK", "BUS", "BICYCLE", "PEDESTRIAN") | |
| RUNTIME_PARAMS = {"score_threshold": (float, 0.0, 1.0, 0.35)} | |
| DEVICE_DEFAULTS = {"num_command_queues": 1, "trace_region_size": 64 << 20} | |
| def _build(self): ... # weights -> device tensors, TraceRunner + variants (no capture) | |
| def _warm_one(self, v): ... # capture the variant's traces (TraceRunner.capture) | |
| def _prepare(self, **kw): ... # host pre-processing -> host buffers | |
| def _forward(self, prep): ... # upload + replay + read (under the model lock) | |
| def _postprocess(self, raw, prep, params): ... # -> an ttaw.outputs object | |
| def _release(self): ... # TraceRunner.release() | |
| Contract: ``from_pretrained`` resolves the pinned weights *before* opening the chip (a Hub problem must not claim | |
| the device), opens it (ETH dispatch, 12x10), builds, and warms every default variant, so the first call is as fast | |
| as the next. Calls are serialized by a re-entrant lock; ``close()`` is idempotent and also runs at interpreter exit. | |
| Importing this module has no side effects (no ttnn, no torch, no network). | |
| """ | |
| from __future__ import annotations | |
| import atexit | |
| import logging | |
| import math | |
| import os | |
| import threading | |
| import time | |
| import weakref | |
| from pathlib import Path | |
| from typing import Any, ClassVar, Dict, List, Mapping, Optional, Sequence, Tuple | |
| import numpy as np | |
| from . import io as tio | |
| from .device import DeviceConfig, close_device, describe_device | |
| __all__ = ["ModelBase", "resolve_weights"] | |
| log = logging.getLogger(__name__) | |
| # the input keyword arguments of ModelBase.__call__ (never runtime params) | |
| _CALL_INPUTS = ("points", "sweeps", "images", "calibration", "stream", "inputs") | |
| def resolve_weights(model_id: str, *, revision: Optional[str] = None, allow_patterns: Optional[Sequence[str]] = None, | |
| weights_dir: Optional[str] = None, env_var: Optional[str] = None) -> Path: | |
| """Local directory holding the weights files. | |
| Order: ``weights_dir`` > ``$<env_var>`` > ``model_id`` if it is a local directory > the Hugging Face snapshot of | |
| ``model_id`` at ``revision`` restricted to ``allow_patterns``. If the Hub is unreachable the cached snapshot is | |
| used (``local_files_only=True``). Never needs a token for public repos; never uploads.""" | |
| for cand in (weights_dir, os.environ.get(env_var) if env_var else None): | |
| if cand: | |
| p = Path(cand).expanduser() | |
| if not p.is_dir(): | |
| raise FileNotFoundError(f"weights directory {p} does not exist") | |
| return p | |
| if Path(model_id).expanduser().is_dir(): | |
| return Path(model_id).expanduser() | |
| from huggingface_hub import snapshot_download | |
| patterns = list(allow_patterns) if allow_patterns else None | |
| try: | |
| return Path(snapshot_download(model_id, revision=revision, allow_patterns=patterns)) | |
| except Exception as e: # noqa: BLE001 -- offline or rate limited: fall back to the local cache | |
| log.warning("Hub not reachable for %s@%s (%s: %s); trying the local HF cache", model_id, revision, | |
| type(e).__name__, e) | |
| return Path(snapshot_download(model_id, revision=revision, allow_patterns=patterns, local_files_only=True)) | |
| _LIVE: "weakref.WeakSet" = weakref.WeakSet() | |
| def _close_all() -> None: | |
| """A model that was not closed is closed when Python exits, so the chip is left clean.""" | |
| for model in list(_LIVE): | |
| try: | |
| model.close() | |
| except Exception: # noqa: BLE001 -- best effort at exit | |
| pass | |
| class ModelBase: | |
| """Base class of every bundle's model class. Create instances with :meth:`from_pretrained`.""" | |
| # ---- identity (override in the bundle) --------------------------------------------------------------------- | |
| MODEL_NAME: ClassVar[str] = "model" | |
| ENV_PREFIX: ClassVar[str] = "MODEL" # <ENV>_WEIGHTS_DIR, <ENV>_DISPATCH, <ENV>_NUM_CQS, ... | |
| DEFAULT_REPO: ClassVar[str] = "" | |
| DEFAULT_TAG: ClassVar[Optional[str]] = None | |
| DEFAULT_REVISION: ClassVar[Optional[str]] = None # the commit DEFAULT_TAG points to (tags can move) | |
| ALLOW_PATTERNS: ClassVar[Optional[Sequence[str]]] = None | |
| WEIGHTS_LICENSE: ClassVar[str] = "Apache-2.0 (AutowareFoundation model card)" | |
| VARIANTS: ClassVar[Sequence[str]] = ("default",) | |
| DEFAULT_VARIANT: ClassVar[str] = "default" | |
| INPUT_KIND: ClassVar[str] = "lidar" # lidar | camera | multicam | lidar+multicam | planner | |
| CAMERA_ORDER: ClassVar[Sequence[str]] = () | |
| REQUIRE_CALIBRATION: ClassVar[Optional[bool]] = None # None: True for multicam / lidar+multicam | |
| POINT_FIELDS: ClassVar[Sequence[str]] = tio.DEFAULT_POINT_FIELDS | |
| LABELS: ClassVar[Sequence[str]] = () | |
| # name -> (type, min, max, default): host-side knobs only (BUNDLE_CONVENTIONS.md section 9); min / max may be | |
| # None for an open side | |
| RUNTIME_PARAMS: ClassVar[Mapping[str, Tuple[type, Any, Any, Any]]] = {} | |
| # model-specific input keyword arguments passed to _prepare (not runtime params), e.g. PointPainting ("rois",): | |
| # the server's ServerSpec.decode_extra adds them to the call kwargs | |
| EXTRA_INPUTS: ClassVar[Sequence[str]] = () | |
| # planner-style named inputs: name -> (shape with None for free dims, dtype); when set, ``inputs=`` is decoded | |
| # and checked with io.load_named_arrays on every call (API and server alike) | |
| INPUT_SCHEMA: ClassVar[Optional[Mapping[str, Tuple[Sequence[Optional[int]], Any]]]] = None | |
| # defaults for ttaw.device.DeviceConfig; <ENV>_* environment variables override them | |
| DEVICE_DEFAULTS: ClassVar[Mapping[str, Any]] = {} | |
| def __init__(self, *_args: Any, **_kwargs: Any): | |
| raise TypeError(f"use {type(self).__name__}.from_pretrained(...)") | |
| # ---- weights / device ------------------------------------------------------------------------------------ | |
| def requires_calibration(cls) -> bool: | |
| if cls.REQUIRE_CALIBRATION is not None: | |
| return bool(cls.REQUIRE_CALIBRATION) | |
| return cls.INPUT_KIND in ("multicam", "lidar+multicam") | |
| def resolve_weights(cls, model_id: Optional[str] = None, revision: Optional[str] = None, | |
| weights_dir: Optional[str] = None) -> Path: | |
| """:func:`resolve_weights` with this model's defaults (pinned revision for the default repo).""" | |
| model_id = model_id or cls.DEFAULT_REPO | |
| rev = revision or (cls.DEFAULT_REVISION if model_id == cls.DEFAULT_REPO else None) | |
| return resolve_weights(model_id, revision=rev, allow_patterns=cls.ALLOW_PATTERNS, weights_dir=weights_dir, | |
| env_var=f"{cls.ENV_PREFIX}_WEIGHTS_DIR") | |
| def device_config(cls, **overrides: Any) -> DeviceConfig: | |
| """``DEVICE_DEFAULTS`` < ``<ENV>_*`` / ``TT_DEVICE_ID`` environment < explicit non-None ``overrides`` | |
| (``device_id``, ``dispatch``, ``num_command_queues``, ...).""" | |
| base = DeviceConfig.from_env(cls.ENV_PREFIX, **dict(cls.DEVICE_DEFAULTS)) | |
| values = {k: getattr(base, k) for k in base.__dataclass_fields__} | |
| values.update({k: v for k, v in overrides.items() if v is not None}) | |
| return DeviceConfig(**values) | |
| def from_pretrained(cls, model_id: Optional[str] = None, *, revision: Optional[str] = None, | |
| variant: Optional[str] = None, device_id: Optional[int] = None, device: Any = None, | |
| dispatch: Optional[str] = None, num_command_queues: Optional[int] = None, | |
| weights_dir: Optional[str] = None, warmup_variants: Any = "default", verbose: bool = False, | |
| **compile_params: Any) -> "ModelBase": | |
| """Resolve weights, open the chip (unless ``device`` is given; ``close()`` then leaves it open), build the | |
| graph and capture the traces of ``warmup_variants`` (``"default"``, ``"none"`` or a list). | |
| ``variant``: one of ``VARIANTS`` (load-time; default ``<ENV>_VARIANT`` or ``DEFAULT_VARIANT``). | |
| ``device_id``: default ``TT_DEVICE_ID`` or 0. ``dispatch``: ``"eth"`` / ``"worker"`` / ``"auto"`` (default | |
| ``<ENV>_DISPATCH`` or eth). ``num_command_queues``: default ``<ENV>_NUM_CQS`` or ``DEVICE_DEFAULTS``. | |
| ``compile_params``: shape-defining options, validated in ``_build``.""" | |
| variant = variant or os.environ.get(f"{cls.ENV_PREFIX}_VARIANT") or cls.DEFAULT_VARIANT | |
| if variant not in cls.VARIANTS: | |
| raise ValueError(f"variant={variant!r}: expected one of {list(cls.VARIANTS)}") | |
| clash = sorted(set(cls.RUNTIME_PARAMS) & (set(_CALL_INPUTS) | set(cls.EXTRA_INPUTS))) | |
| if clash: | |
| raise TypeError(f"{cls.__name__}: RUNTIME_PARAMS {clash} clash with input keyword arguments") | |
| self = cls.__new__(cls) | |
| self._lock = threading.RLock() | |
| self._closed = False | |
| self._owns_device = device is None | |
| self.device = None | |
| self.variant = variant | |
| self.compile_params = dict(compile_params) | |
| self.verbose = verbose | |
| self.warmup_ms: Dict[str, float] = {} | |
| self.warm_variants: List[Any] = [] | |
| self.device_info: Dict[str, Any] = {} | |
| model_id = model_id or cls.DEFAULT_REPO | |
| rev = revision or (cls.DEFAULT_REVISION if model_id == cls.DEFAULT_REPO else None) | |
| t0 = time.perf_counter() | |
| log.info("Loading weights %s@%s (variant %s)", model_id, rev, variant) | |
| self.weights_path = cls.resolve_weights(model_id, revision, weights_dir) | |
| self.weights = {"repo": model_id, "tag": cls.DEFAULT_TAG if model_id == cls.DEFAULT_REPO else None, | |
| "revision": rev, "path": str(self.weights_path), "license": cls.WEIGHTS_LICENSE} | |
| self.warmup_ms["weights"] = (time.perf_counter() - t0) * 1e3 | |
| try: | |
| if device is None: | |
| cfg = cls.device_config(device_id=device_id, dispatch=dispatch, num_command_queues=num_command_queues) | |
| log.info("Opening device %d (dispatch %s, %d CQ)", cfg.device_id, cfg.dispatch, cfg.num_command_queues) | |
| device = cfg.open() | |
| device_id = cfg.device_id | |
| self.device = device | |
| self.device_info = describe_device(device, device_id) | |
| t1 = time.perf_counter() | |
| self._build() | |
| self.warmup_ms["build"] = (time.perf_counter() - t1) * 1e3 | |
| self.warmup(warmup_variants) | |
| except BaseException: | |
| self.close() | |
| raise | |
| self.warmup_ms["total"] = (time.perf_counter() - t0) * 1e3 | |
| _LIVE.add(self) | |
| return self | |
| # ---- warm-up --------------------------------------------------------------------------------------------- | |
| def default_warmup_variants(self) -> List[Any]: | |
| """The trace variants captured by default (e.g. pillar-count buckets). Override per port.""" | |
| return [{"variant": self.variant}] | |
| def warmup(self, variants: Any = "default") -> Dict[str, Any]: | |
| """Compile + capture each variant (idempotent). Returns ``{"ms": ..., "variants": [...]}``.""" | |
| self._check_open() | |
| if variants == "default": | |
| todo = self.default_warmup_variants() | |
| elif variants in (None, "none"): | |
| todo = [] | |
| else: | |
| todo = list(variants) | |
| t0 = time.perf_counter() | |
| with self._lock: | |
| for v in todo: | |
| if v in self.warm_variants: | |
| continue | |
| log.info("Warming up: capturing trace for %s", v) | |
| self._warm_one(v) | |
| self.warm_variants.append(v) | |
| ms = (time.perf_counter() - t0) * 1e3 | |
| self.warmup_ms["warmup"] = self.warmup_ms.get("warmup", 0.0) + ms | |
| log.info("Warmup complete: %d variant(s) in %.0f ms", len(self.warm_variants), ms) | |
| return {"ms": ms, "variants": list(self.warm_variants)} | |
| # ---- inference ------------------------------------------------------------------------------------------- | |
| def validate_params(cls, params: Optional[Mapping[str, Any]]) -> Dict[str, Any]: | |
| """Per-request knobs -> typed values with defaults. Unknown names, values outside ``[min, max]`` (either | |
| bound may be None), non-booleans for ``bool``, and booleans or non-integral numbers for ``int`` raise | |
| ``InputError``; numeric strings are accepted for numbers. ``None`` is accepted for a param whose default | |
| is ``None`` (it means "not set"), so ``validate_params(validate_params(p)) == validate_params(p)``.""" | |
| out = {k: spec[3] for k, spec in cls.RUNTIME_PARAMS.items()} | |
| for k, v in (params or {}).items(): | |
| if k not in cls.RUNTIME_PARAMS: | |
| raise tio.InputError(f"unknown parameter {k!r}; allowed: {sorted(cls.RUNTIME_PARAMS)}") | |
| typ, lo, hi, default = cls.RUNTIME_PARAMS[k] | |
| if v is None and default is None: # "not set" (JSON null) = the default, so validation is idempotent: | |
| out[k] = None # the server validates, then model(**params) validates again | |
| continue | |
| if typ is bool and not isinstance(v, bool): | |
| raise tio.InputError(f"parameter {k!r} must be a boolean") | |
| if typ in (int, float) and isinstance(v, bool): | |
| raise tio.InputError(f"parameter {k!r} must be {typ.__name__}, not a boolean") | |
| try: | |
| if typ is int and not isinstance(v, int): | |
| number = float(v) | |
| if not number.is_integer(): | |
| raise ValueError(v) | |
| v = int(number) | |
| else: | |
| v = typ(v) | |
| except (TypeError, ValueError, OverflowError): | |
| raise tio.InputError(f"parameter {k!r} must be {typ.__name__}") from None | |
| if typ is float and not math.isfinite(v): | |
| raise tio.InputError(f"parameter {k!r} must be finite") | |
| if (lo is not None and v < lo) or (hi is not None and v > hi): | |
| raise tio.InputError(f"parameter {k!r}={v} is outside [{lo}, {hi}]") | |
| out[k] = v | |
| return out | |
| def __call__(self, points: Any = None, *, sweeps: Optional[Sequence] = None, images: Any = None, | |
| calibration: Optional[Mapping] = None, stream: Optional[Mapping] = None, | |
| inputs: Any = None, **params: Any): | |
| """One frame (batch 1). ``points``: path / bytes / (N, C) array / PointCloud / JSON envelope. | |
| ``sweeps``: ``[{"points", "time_lag_s", "T_current_from_sweep"}]``. ``images``: list of camera dicts or | |
| CameraImage. ``calibration``: dict (or ``{"preset": name}``, resolved by the server). ``stream``: | |
| ``{"id", "reset", "timestamp_s", "T_world_from_ego"}``. ``inputs``: named arrays (planner; checked | |
| against ``INPUT_SCHEMA`` when the class sets one). Keyword arguments named in ``EXTRA_INPUTS`` go to | |
| ``_prepare``; the others are ``RUNTIME_PARAMS``. Returns the model's output object with ``timing_ms``.""" | |
| self._check_open() | |
| extra = {k: params.pop(k) for k in list(params) if k in self.EXTRA_INPUTS} | |
| p = self.validate_params(params) | |
| if inputs is not None and self.INPUT_SCHEMA is not None: | |
| inputs = tio.load_named_arrays(inputs, self.INPUT_SCHEMA) | |
| with self._lock: | |
| self._check_open() # close() may have run while this call waited for the lock | |
| t0 = time.perf_counter() | |
| prepared = self._prepare(points=points, sweeps=sweeps, images=images, calibration=calibration, | |
| stream=stream, inputs=inputs, **extra) | |
| t1 = time.perf_counter() | |
| raw = self._forward(prepared) | |
| t2 = time.perf_counter() | |
| out = self._postprocess(raw, prepared, p) | |
| t3 = time.perf_counter() | |
| out.timing_ms.update({"preprocess": (t1 - t0) * 1e3, "device": (t2 - t1) * 1e3, | |
| "postprocess": (t3 - t2) * 1e3, "total": (t3 - t0) * 1e3}) | |
| return out | |
| predict = __call__ | |
| # ---- lifetime -------------------------------------------------------------------------------------------- | |
| def extra_info(self) -> Dict[str, Any]: | |
| """Port-specific additions to :attr:`info` (knobs, precision policy, TraceRunner.describe() ...).""" | |
| return {} | |
| def info(self) -> Dict[str, Any]: | |
| cls = type(self) | |
| info = {"model": cls.MODEL_NAME, "variant": getattr(self, "variant", None), | |
| "weights": getattr(self, "weights", None), "device": getattr(self, "device_info", None), | |
| "warm_variants": getattr(self, "warm_variants", []), "warmup_ms": getattr(self, "warmup_ms", {}), | |
| "compile_params": getattr(self, "compile_params", {}), "input_kind": cls.INPUT_KIND, | |
| "point_fields": list(cls.POINT_FIELDS), "camera_order": list(cls.CAMERA_ORDER), | |
| "labels": list(cls.LABELS), "runtime_params": {k: v[3] for k, v in cls.RUNTIME_PARAMS.items()}, | |
| "extra_inputs": list(cls.EXTRA_INPUTS)} | |
| if cls.INPUT_SCHEMA is not None: | |
| info["input_schema"] = {name: {"shape": [None if d is None else int(d) for d in shape], | |
| "dtype": np.dtype(dtype).name} | |
| for name, (shape, dtype) in cls.INPUT_SCHEMA.items()} | |
| if not getattr(self, "_closed", True): | |
| info.update(self.extra_info()) | |
| return info | |
| def closed(self) -> bool: | |
| return getattr(self, "_closed", True) | |
| def _check_open(self) -> None: | |
| if self.closed: | |
| raise RuntimeError(f"{type(self).__name__} is closed") | |
| def close(self) -> None: | |
| """Release traces and device tensors; close the chip if this model opened it. Idempotent.""" | |
| if getattr(self, "_closed", True): | |
| return | |
| self._closed = True | |
| _LIVE.discard(self) | |
| with self._lock: | |
| try: | |
| self._release() | |
| finally: | |
| dev = getattr(self, "device", None) | |
| if dev is not None and self._owns_device: | |
| close_device(dev) | |
| self.device = None | |
| def __enter__(self) -> "ModelBase": | |
| return self | |
| def __exit__(self, *exc: Any) -> None: | |
| self.close() | |
| def __repr__(self) -> str: | |
| state = ("closed" if self.closed | |
| else f"{self.device_info.get('dispatch')} {self.device_info.get('grid')}") | |
| return f"<{type(self).__name__} {self.MODEL_NAME} variant={getattr(self, 'variant', '?')} {state}>" | |
| # ---- port-specific hooks --------------------------------------------------------------------------------- | |
| def _build(self) -> None: | |
| """Load weights from ``self.weights_path``, convert them to device tensors, build the ttnn graph and its | |
| ``TraceRunner`` (inputs, params, states, variants). No capture here.""" | |
| raise NotImplementedError("port-specific: build the ttnn graph") | |
| def _warm_one(self, variant: Any) -> None: | |
| """Warm up and capture the traces of one variant (``TraceRunner.capture``).""" | |
| raise NotImplementedError("port-specific: compile + capture one variant") | |
| def _prepare(self, **kwargs: Any) -> Any: | |
| """Host pre-processing ported from the Autoware node -> host buffers. Raise ``InputError`` for client | |
| mistakes (missing input, wrong shape).""" | |
| raise NotImplementedError("port-specific: host pre-processing") | |
| def _forward(self, prepared: Any) -> Any: | |
| """Upload into the persistent device inputs, replay the matching trace, read the outputs.""" | |
| raise NotImplementedError("port-specific: traced device forward") | |
| def _postprocess(self, raw: Any, prepared: Any, params: Dict[str, Any]) -> Any: | |
| """Host post-processing ported from Autoware -> an output object of :mod:`ttaw.outputs`.""" | |
| raise NotImplementedError("port-specific: host post-processing") | |
| def _release(self) -> None: | |
| """Release traces and persistent device tensors (``TraceRunner.release()``). Also called when | |
| ``from_pretrained`` fails half-way, so tolerate a partially built model (``getattr(self, "runner", None)``).""" | |