# 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`` > ``$`` > ``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() @atexit.register 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" # _WEIGHTS_DIR, _DISPATCH, _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; _* 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 ------------------------------------------------------------------------------------ @classmethod 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") @classmethod 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") @classmethod def device_config(cls, **overrides: Any) -> DeviceConfig: """``DEVICE_DEFAULTS`` < ``_*`` / ``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) @classmethod 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 ``_VARIANT`` or ``DEFAULT_VARIANT``). ``device_id``: default ``TT_DEVICE_ID`` or 0. ``dispatch``: ``"eth"`` / ``"worker"`` / ``"auto"`` (default ``_DISPATCH`` or eth). ``num_command_queues``: default ``_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 ------------------------------------------------------------------------------------------- @classmethod 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 {} @property 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 @property 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)``)."""