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