Download code/models/server/app.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 38.3 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/server/app.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/server/app.py
-
curl -L -o app.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/server/app.py
38.3 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """ASGI serving app for SuperPoint on one Tenstorrent Blackhole chip. | |
| Served by tt-model-manager as ``kind: tt-dit-server``:: | |
| runtime: | |
| app: models.server.app:app | |
| uvicorn runs this module. Everything that touches the device, the network or | |
| the weights happens inside the ASGI **lifespan**: uvicorn's ``Application | |
| startup complete`` -- the line ``tt-model serve`` waits for -- therefore means | |
| the chip is claimed, the weights are loaded and the kernels are compiled. | |
| Serving path (default, ``TT_FUSED`` unset or 1; device-validated 2026-09-13, see | |
| DEVICE_VALIDATION.md "Results"): fixed 480x640 input, the whole device graph as ONE | |
| metal trace -- 64-byte-page input upload, encoder + heads (pure ttnn, on-device | |
| softmax), ``rms_norm`` descriptor L2-norm, the standard-op device NMS (radius 4, the | |
| trace default), row-major outputs -- captured during warm-up (compile pass -> capture | |
| -> one replay, all before READY) and replayed per request; the host then runs | |
| threshold/border/top-k/grid_sample only. Since 2026-10-03 (OPT_REPORT.md round 1) the default | |
| stages also put the keypoint list + bilinear descriptor sampling (``kpc``) and the uint8 -> bf16 | |
| input conversion (``u8``) in the trace: a request uploads the 8-bit R plane (307 KB) and reads back | |
| an 8 KB keypoint header + the sampled descriptor rows (``_infer_fused``), bit-identical to the | |
| host post-processing. Requests with ``nms_radius != 4`` take the | |
| host NMS from the traced scores (same output, slower). No custom kernel, one image | |
| per request. ``TT_FUSED_STAGES`` / ``SP_TRACE_REGION`` are the device A/B knobs | |
| (models/tt/fused_host.py). | |
| ``TT_FUSED=0`` (read once in the lifespan) restores the legacy path the port validated | |
| first (models/tests/test_superpoint.py, models/tt/postprocess.py): untraced pure-ttnn | |
| device forward, host fold + NMS -- byte-identical to the 2026-09-12 shipped server. | |
| Environment (read in the lifespan, never at import): | |
| HF_MODEL weights repo id (default magic-leap-community/superpoint) | |
| TT_WEIGHTS_REVISION commit sha to load (default: the repo's default branch) | |
| SP_WEIGHTS_DIR local directory with config.json + model.safetensors | |
| (overrides HF_MODEL / TT_WEIGHTS_REVISION; offline/host use) | |
| TT_MESH_SHAPE "1x1" (also "(1, 1)" / "1,1"); anything else is refused | |
| TT_DEVICE_ID chip to open (default 0) | |
| SP_DISPATCH unset/auto = ETH dispatch + 1 command queue + 12x10 grid when the tt-metal | |
| ETH-dispatch patch is present (the p150 target), else Tensix dispatch with a | |
| warning; "eth" forces ETH; "worker" = Tensix dispatch (explicit opt-in; | |
| 11x10 on a p150, 12x10 only on a Galaxy chip) | |
| TT_FUSED unset/1 = traced fused path (default); "0" = legacy untraced path | |
| TT_FUSED_STAGES fused stages (default all: wide,nms,rms,rm,l1,nmsk,kpc,u8,rsz); device A/B only | |
| SP_RSZ_PRECOMPILE rsz stage: source sizes (WxH, comma list) whose device-resize variant is captured | |
| before READY (default 1920x1080; others on first use, ~5-12 ms once) | |
| SP_TRACE_REGION trace_region_size bytes on the fused path (default 32 MiB) | |
| SP_NMS_RADII_PRECOMPILE comma list of nms_radius values (1..8) whose device NMS variant is | |
| captured before READY (default: none, each is built on first use) | |
| I/O: one base64 PNG/JPEG -> keypoints (original-image pixel coordinates), | |
| scores, and optionally 256-d L2-normalised descriptors (float16 NPZ, base64). | |
| """ | |
| from __future__ import annotations | |
| import base64 | |
| import binascii | |
| import io | |
| import logging | |
| import os | |
| import re | |
| import sys | |
| import threading | |
| import time | |
| from contextlib import asynccontextmanager | |
| from typing import Any, Dict, Tuple | |
| import numpy as np | |
| import pydantic_core | |
| import torch | |
| from fastapi import FastAPI, HTTPException, Query, Request | |
| from fastapi.concurrency import run_in_threadpool | |
| from fastapi.exceptions import RequestValidationError | |
| from fastapi.responses import JSONResponse, Response | |
| from PIL import Image | |
| from pydantic import BaseModel, Field | |
| # torch-only; the ttnn-importing port module is imported lazily in the lifespan. | |
| from ..tt import device_open as _devopen | |
| from ..tt import fused_host as _fused | |
| from ..tt import postprocess as _post | |
| LOG = logging.getLogger("superpoint.server") | |
| MODEL_NAME = "superpoint-p150" | |
| TASK = "keypoint-detection" | |
| DEFAULT_WEIGHTS_REPO = "magic-leap-community/superpoint" | |
| SOURCE_REPO = "https://github.com/changh95/tt-superpoint" | |
| SOURCE_COMMIT = "e1eab66e29ff424bc9af6b1118671d9bc08e899e" | |
| LICENSE_NOTE = ( | |
| "Weights: Magic Leap SuperPoint licence -- academic or non-profit organisation " | |
| "NONCOMMERCIAL research use only (see https://huggingface.co/magic-leap-community/superpoint). " | |
| "Port code: Apache-2.0 headers, distributed under the same upstream terms." | |
| ) | |
| # The port's validated canonical input (HF SuperPointImageProcessor default size). | |
| INPUT_HEIGHT, INPUT_WIDTH = 480, 640 | |
| # Device-open kwargs of the port's own untraced end-to-end script (models/visualize.py). | |
| L1_SMALL_SIZE = 32 * 1024 | |
| # Hard cap on max_keypoints: one candidate per pixel of the canonical frame. | |
| MAX_KEYPOINTS_CAP = INPUT_HEIGHT * INPUT_WIDTH | |
| STATE: Dict[str, Any] = {"ready": False} | |
| LOCK = threading.Lock() # every device-touching call goes through here | |
| # --------------------------------------------------------------------------- config | |
| def _setup_logging() -> None: | |
| """Make our INFO lines show up on uvicorn's stdout (drives tt-model's boot checklist).""" | |
| if not LOG.handlers: | |
| handler = logging.StreamHandler() | |
| handler.setFormatter(logging.Formatter("%(levelname)s: [superpoint] %(message)s")) | |
| LOG.addHandler(handler) | |
| LOG.setLevel(logging.INFO) | |
| LOG.propagate = False | |
| def _parse_mesh_shape(raw: str) -> Tuple[int, int]: | |
| """Accept "1x1", "(1, 1)", "1,1", "[1, 1]". Anything else -> RuntimeError.""" | |
| nums = re.findall(r"\d+", raw or "") | |
| if len(nums) != 2: | |
| raise RuntimeError( | |
| f"TT_MESH_SHAPE={raw!r} is not a mesh shape; expected 'RxC' such as '1x1'" | |
| ) | |
| return int(nums[0]), int(nums[1]) | |
| def _config_from_env() -> Dict[str, Any]: | |
| rows, cols = _parse_mesh_shape(os.environ.get("TT_MESH_SHAPE", "1x1")) | |
| if (rows, cols) != (1, 1): | |
| raise RuntimeError( | |
| f"TT_MESH_SHAPE={rows}x{cols} is a multi-chip mesh; this port runs on a single " | |
| "chip (mesh_device: P150). Serve it with hardware p150 / mesh 1x1." | |
| ) | |
| weights_dir = os.environ.get("SP_WEIGHTS_DIR") or None | |
| fused = _fused.fused_enabled() # TT_FUSED, read once here | |
| return { | |
| "mesh_shape": (rows, cols), | |
| "device_id": int(os.environ.get("TT_DEVICE_ID", "0")), | |
| "weights_repo": os.environ.get("HF_MODEL") or DEFAULT_WEIGHTS_REPO, | |
| "weights_revision": os.environ.get("TT_WEIGHTS_REVISION") or None, | |
| "weights_dir": weights_dir, | |
| "fused": fused, | |
| "fused_stages": sorted(_fused.fused_stages()) if fused else [], | |
| "trace_region_size": _fused.trace_region_size() if fused else 0, | |
| "dispatch": os.environ.get("SP_DISPATCH", "auto") or "auto", | |
| } | |
| # --------------------------------------------------------------------------- weights | |
| def _load_reference(cfg: Dict[str, Any]): | |
| """Load the fp32 HF reference whose weights the tt-nn model wraps. | |
| Uses ``transformers.SuperPointForKeypointDetection.from_pretrained`` with the | |
| pinned ``revision`` so the sha ``tt-model serve`` pre-downloaded into the HF | |
| cache (mounted at /hf) is what gets loaded: a sha-pinned snapshot has no | |
| ``refs/main``, so resolving ``main`` would need the network and could pick | |
| different weights. If the Hub is unreachable but the pinned snapshot is | |
| cached, the ``local_files_only`` retry still boots. | |
| """ | |
| from transformers import SuperPointForKeypointDetection | |
| if cfg["weights_dir"]: | |
| src = cfg["weights_dir"] | |
| if not os.path.isdir(src): | |
| raise RuntimeError(f"SP_WEIGHTS_DIR={src!r} is not a directory") | |
| LOG.info("Loading weights from local directory %s", src) | |
| model = SuperPointForKeypointDetection.from_pretrained(src) | |
| else: | |
| repo, rev = cfg["weights_repo"], cfg["weights_revision"] | |
| LOG.info("Loading weights %s @ %s", repo, rev or "default branch") | |
| try: | |
| model = SuperPointForKeypointDetection.from_pretrained(repo, revision=rev) | |
| except Exception as e: # network down, pinned snapshot cached -> use it | |
| if rev is None: | |
| raise | |
| LOG.warning("Hub resolution failed (%s: %s); retrying from the local HF cache", type(e).__name__, e) | |
| model = SuperPointForKeypointDetection.from_pretrained(repo, revision=rev, local_files_only=True) | |
| model.eval() | |
| return model | |
| # --------------------------------------------------------------------------- lifespan | |
| async def lifespan(_app: FastAPI): | |
| _setup_logging() | |
| torch.set_grad_enabled(False) | |
| cfg = _config_from_env() | |
| STATE["cfg"] = cfg | |
| t0 = time.perf_counter() | |
| torch_model = _load_reference(cfg) | |
| STATE["model_config"] = { | |
| "nms_radius": int(torch_model.config.nms_radius), | |
| "keypoint_threshold": float(torch_model.config.keypoint_threshold), | |
| "max_keypoints": int(torch_model.config.max_keypoints), | |
| "border_removal_distance": int(torch_model.config.border_removal_distance), | |
| } | |
| t_weights = time.perf_counter() - t0 | |
| LOG.info("Loading weights done in %.1fs", t_weights) | |
| import ttnn # in the image and the tt-metal venv; deliberately not at import time | |
| from ..tt.superpoint_ttnn import TtSuperPoint | |
| # Legacy (TT_FUSED=0): exactly CreateDevice(device_id=..., l1_small_size=...). The fused | |
| # path adds the trace region (ttnn's default 0 makes begin_trace_capture impossible). | |
| open_kwargs = _fused.device_open_kwargs(cfg["device_id"], L1_SMALL_SIZE, cfg["fused"], cfg["trace_region_size"]) | |
| # Dispatch (OPT_REPORT.md "p150 ETH-dispatch compliance 2026-10-05"): SP_DISPATCH unset/auto = ETH | |
| # dispatch + 1 command queue + 12x10 compute grid when the tt-metal ETH-dispatch patch is present | |
| # (the p150 target), else the Tensix dispatch with a warning; "worker" = explicit Tensix opt-in. | |
| dispatch = _devopen.resolve_dispatch(cfg["dispatch"]) | |
| LOG.info( | |
| "Opening device %d (dispatch=%s, %s, mesh %dx%d)", | |
| cfg["device_id"], | |
| dispatch, | |
| ", ".join(f"{k}={v}" for k, v in open_kwargs.items() if k != "device_id"), | |
| *cfg["mesh_shape"], | |
| ) | |
| dev_id = open_kwargs.pop("device_id") | |
| device = _devopen.open_ttnn_device(dev_id, dispatch=dispatch, **open_kwargs) | |
| g = device.compute_with_storage_grid_size() | |
| cfg["dispatch"] = dispatch | |
| cfg["compute_grid"] = f"{g.x}x{g.y}" | |
| cfg["num_command_queues"] = 1 | |
| LOG.info("Device open: dispatch=%s, 1 command queue, compute grid %dx%d", dispatch, g.x, g.y) | |
| if dispatch == "worker": | |
| LOG.warning( | |
| "Tensix (worker) dispatch: explicit opt-in. On a p150 this leaves an 11x10 compute grid; " | |
| "the published numbers use ETH dispatch (12x10). Unset SP_DISPATCH for the default." | |
| ) | |
| STATE["device"] = device | |
| STATE["ttnn"] = ttnn | |
| try: | |
| model = TtSuperPoint( | |
| torch_model, device, input_height=INPUT_HEIGHT, input_width=INPUT_WIDTH, fused=cfg["fused"] | |
| ) | |
| tt_in = model.allocate_input(batch_size=1) | |
| STATE["model"] = model | |
| STATE["tt_in"] = tt_in | |
| STATE["fused"] = bool(model.fused) | |
| del torch_model # the tt-nn model holds its own copies of the weights | |
| dummy = torch.zeros(1, 3, INPUT_HEIGHT, INPUT_WIDTH, dtype=torch.float32) | |
| if model.fused: | |
| _warmup_fused(model, tt_in, dummy, ttnn, device) | |
| else: | |
| # Warm up: the first forward JIT-compiles every kernel and converts the | |
| # conv weights to their device layout; the second one measures steady state. | |
| LOG.info("Warming up (compiling kernels on a %dx%d dummy frame) ...", INPUT_HEIGHT, INPUT_WIDTH) | |
| timings = [] | |
| with LOCK, torch.inference_mode(): | |
| for _ in range(2): | |
| t1 = time.perf_counter() | |
| scores, desc = model.run_untraced(tt_in, dummy) | |
| ttnn.synchronize_device(device) | |
| timings.append((time.perf_counter() - t1) * 1000.0) | |
| if not (torch.isfinite(scores).all() and torch.isfinite(desc).all()): | |
| raise RuntimeError("warm-up forward produced non-finite outputs") | |
| STATE["warmup_ms"] = {"first_forward": round(timings[0], 1), "second_forward": round(timings[1], 1)} | |
| LOG.info("Warmup complete: first forward %.0f ms (compile), second %.0f ms", timings[0], timings[1]) | |
| STATE["ready"] = True | |
| yield | |
| finally: | |
| STATE["ready"] = False | |
| _shutdown() | |
| def _fused_result_finite(res) -> bool: | |
| ok = bool(torch.isfinite(res.descriptors_nchw).all()) | |
| if res.nms_map is not None: | |
| ok = ok and bool(torch.isfinite(res.nms_map).all()) | |
| if res.scores_nchw is not None: | |
| ok = ok and bool(torch.isfinite(res.scores_nchw).all()) | |
| return ok | |
| def _warmup_fused(model, tt_in, dummy: torch.Tensor, ttnn, device) -> None: | |
| """TT_FUSED warm-up contract: eager compile pass -> trace capture -> one traced replay, | |
| all before READY. A knob-on server that could not capture its trace must not come up.""" | |
| stages = sorted(model.fused_stages) | |
| LOG.info( | |
| "Warming up TT_FUSED path (stages %s, traced nms_radius %d) on a %dx%d dummy frame ...", | |
| ",".join(stages) or "trace-only", model.nms_radius_traced, INPUT_HEIGHT, INPUT_WIDTH, | |
| ) | |
| timings = {} | |
| with LOCK, torch.inference_mode(): | |
| t1 = time.perf_counter() | |
| res = model.run_fused(tt_in, dummy) # eager: compiles kernels, prepares conv weights | |
| ttnn.synchronize_device(device) | |
| timings["compile_forward"] = (time.perf_counter() - t1) * 1000.0 | |
| if not _fused_result_finite(res): | |
| raise RuntimeError("warm-up (eager fused) forward produced non-finite outputs") | |
| t1 = time.perf_counter() | |
| model.capture_trace(tt_in, b=1) | |
| timings["trace_capture"] = (time.perf_counter() - t1) * 1000.0 | |
| t1 = time.perf_counter() | |
| res = model.run_fused(tt_in, dummy) # first replay | |
| ttnn.synchronize_device(device) | |
| timings["traced_forward"] = (time.perf_counter() - t1) * 1000.0 | |
| if not _fused_result_finite(res): | |
| raise RuntimeError("warm-up (traced) forward produced non-finite outputs") | |
| # Optional: precompile per-radius device NMS variants before READY (otherwise built on first use). | |
| pre = [int(v) for v in os.environ.get("SP_NMS_RADII_PRECOMPILE", "").split(",") if v.strip()] | |
| with LOCK, torch.inference_mode(): | |
| for r in pre: | |
| if model.supports_device_nms_radius(r): | |
| model._variant(r) | |
| # Device resize: compile the kernel and capture the per-size variants listed in SP_RSZ_PRECOMPILE | |
| # (default 1920x1080) before READY; other source sizes are captured on first use. | |
| if model.device_resize: | |
| t1 = time.perf_counter() | |
| for wh in os.environ.get("SP_RSZ_PRECOMPILE", "1920x1080").split(","): | |
| if wh.strip(): | |
| w, h = (int(v) for v in wh.lower().split("x")) | |
| if model.supports_device_resize(w, h): | |
| _infer_fused(model, tt_in, np.zeros((h, w), dtype=np.uint8), max_keypoints=1024, | |
| keypoint_threshold=float(model.keypoint_threshold), nms_radius=model.nms_radius_traced, | |
| return_descriptors=True, border=int(model.border_removal_distance)) | |
| timings["resize_variants"] = (time.perf_counter() - t1) * 1000.0 | |
| # First request-path call (allocates the persistent readback buffers of the kpc path). | |
| z8 = np.zeros((INPUT_HEIGHT, INPUT_WIDTH), dtype=np.uint8) | |
| t1 = time.perf_counter() | |
| _infer_fused(model, tt_in, z8, max_keypoints=1024, keypoint_threshold=float(model.keypoint_threshold), | |
| nms_radius=model.nms_radius_traced, return_descriptors=True, border=int(model.border_removal_distance)) | |
| timings["request_path"] = (time.perf_counter() - t1) * 1000.0 | |
| if model.trace_id is None: | |
| raise RuntimeError("warm-up did not capture the metal trace (TtSuperPoint.trace_id is None)") | |
| STATE["warmup_ms"] = {k: round(v, 1) for k, v in timings.items()} | |
| LOG.info( | |
| "Warmup complete: compile forward %.0f ms, trace capture %.0f ms, traced forward %.1f ms", | |
| timings["compile_forward"], timings["trace_capture"], timings["traced_forward"], | |
| ) | |
| def _shutdown() -> None: | |
| ttnn = STATE.pop("ttnn", None) | |
| device = STATE.pop("device", None) | |
| tt_in = STATE.pop("tt_in", None) | |
| model = STATE.pop("model", None) | |
| STATE.pop("fused", None) | |
| if ttnn is None or device is None: | |
| return | |
| try: | |
| ttnn.synchronize_device(device) | |
| if model is not None and getattr(model, "fused", False): | |
| model.release() # trace + resident fused outputs + gamma | |
| if tt_in is not None: | |
| ttnn.deallocate(tt_in) | |
| del model # drops the device-resident conv weights/biases | |
| except Exception as e: # never let teardown mask the real error | |
| LOG.warning("device tensor release failed: %s", e) | |
| LOG.info("Closing device") | |
| ttnn.close_device(device) | |
| app = FastAPI(title="SuperPoint on Blackhole", lifespan=lifespan) | |
| async def _bad_request(_request: Request, exc: RequestValidationError) -> JSONResponse: | |
| """Malformed request bodies are 400 (the contract), not FastAPI's default 422.""" | |
| return JSONResponse(status_code=400, content={"detail": exc.errors()}) | |
| # --------------------------------------------------------------------------- schema | |
| class PredictRequest(BaseModel): | |
| """One image -> keypoints. ``image`` is a base64-encoded PNG or JPEG.""" | |
| image: str = Field(..., description="base64 PNG/JPEG (RGB or grayscale); resized to 640x480 server-side") | |
| max_keypoints: int = Field( | |
| 1024, ge=-1, le=MAX_KEYPOINTS_CAP, | |
| description="keep the top-k by score; -1 keeps every keypoint above the threshold", | |
| ) | |
| keypoint_threshold: float = Field(0.005, ge=0.0, le=1.0, description="minimum post-NMS score") | |
| nms_radius: int = Field(4, ge=0, le=32, description="single-pass NMS radius in canonical-frame pixels (0 = off)") | |
| return_descriptors: bool = Field(True, description="include 256-d descriptors as a float16 NPZ (base64)") | |
| # --------------------------------------------------------------------------- helpers | |
| _B64_ALPHABET = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/" | |
| _A2B_STRICT = sys.version_info >= (3, 11) | |
| def _b64decode_strict(b64: str) -> bytes: | |
| """``base64.b64decode(b64, validate=True)`` with the same accept / reject set, ~1.4x faster on the | |
| 400 KB request strings (0.75 vs 1.0 ms here): the stdlib check is a regex fullmatch of | |
| ``[A-Za-z0-9+/]*={0,2}``; this is the same predicate as a C-speed ``bytes.translate`` delete of the | |
| alphabet after stripping at most two trailing '=' (models/tests/test_fused_host.py).""" | |
| b = b64.encode("ascii") if isinstance(b64, str) else bytes(b64) | |
| if _A2B_STRICT: | |
| # Python >= 3.11 (the image is 3.12): base64.b64decode(validate=True) IS | |
| # a2b_base64(s, strict_mode=True) (3.11 fuzz, 300k cases: same bytes, same errors); one C pass. | |
| # The translate branch below matches the 3.10 stdlib, which still accepted misplaced '=' padding | |
| # ("=QUJD", "QUJD=") that 3.11+ rejects. | |
| return binascii.a2b_base64(b, strict_mode=True) | |
| t = b.rstrip(b"=") | |
| if len(b) - len(t) > 2 or t.translate(None, _B64_ALPHABET): | |
| raise binascii.Error("Non-base64 digit found") | |
| return binascii.a2b_base64(b) | |
| def _decode_image(b64: str) -> Image.Image: | |
| try: | |
| raw = _b64decode_strict(b64) | |
| except Exception as e: | |
| raise HTTPException(status_code=400, detail=f"bad image: {type(e).__name__}: {e}") from None | |
| return _open_image(raw) | |
| def _open_image(raw: bytes) -> Image.Image: | |
| try: | |
| im = Image.open(io.BytesIO(raw)) | |
| im.load() | |
| return im if im.mode == "RGB" else im.convert("RGB") | |
| except Exception as e: | |
| raise HTTPException(status_code=400, detail=f"bad image: {type(e).__name__}: {e}") from None | |
| def _preprocess(im: Image.Image) -> torch.Tensor: | |
| """PIL RGB -> fp32 (1, 3, 480, 640) in [0, 1]. | |
| Matches the HF ``SuperPointImageProcessor`` defaults the port was validated | |
| against: bilinear resize to 480x640, rescale by 1/255, no grayscale | |
| conversion -- the model then reads channel 0 (R) exactly as | |
| ``SuperPointForKeypointDetection.extract_one_channel_pixel_values`` does. | |
| """ | |
| im = im.resize((INPUT_WIDTH, INPUT_HEIGHT), resample=Image.BILINEAR) | |
| arr = np.asarray(im, dtype=np.float32) / 255.0 # (H, W, 3) | |
| return torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).contiguous() | |
| def _preprocess_r8(im: Image.Image) -> np.ndarray: | |
| """Fused-path preprocess: only channel 0 (R), the one the model reads. PIL's bilinear resize | |
| filters every band independently with the same fixed-point coefficients, so resizing the R | |
| band alone gives exactly the R plane of ``_preprocess`` (before /255); the /255 + bf16 cast | |
| is a table lookup in ``TtSuperPoint.prepare_host_input_u8`` (bit-identical host tensor, | |
| asserted in models/tests/test_fused_host.py). About 3x less resize work than the RGB frame.""" | |
| r = im.getchannel(0).resize((INPUT_WIDTH, INPUT_HEIGHT), resample=Image.BILINEAR) | |
| return np.array(r, dtype=np.uint8) # writable copy (torch.as_tensor on a read-only view warns) | |
| def _r_plane(im: Image.Image) -> np.ndarray: | |
| """``rsz`` stage preprocess: channel 0 (R) at the decoded size; the bilinear resize to 640x480 | |
| runs on device, bit-identical to ``_preprocess_r8``'s Pillow resize (models/tt/resize_r8.py, | |
| asserted in models/tests/test_fused_host.py and test_superpoint.py). Sizes the device kernel | |
| does not cover fall back to the host Pillow resize inside ``TtSuperPoint.prepare_source``.""" | |
| return np.asarray(im.getchannel(0)) | |
| def _infer_fused(model, tt_in, r8: np.ndarray, *, max_keypoints: int, keypoint_threshold: float, | |
| nms_radius: int, return_descriptors: bool, border: int): | |
| """One fused request on the uint8 R plane -> (kp (N,2) xy, scores (N,), desc (N,256) | None, | |
| device_nms). Device lock held for the device part only. | |
| nms_radius == traced radius (4): ``run_fused_keypoints_kpc`` -- one H2D, one trace (network, | |
| NMS, keypoint list, bilinear descriptor sampling), D2H of the keypoint header and of the | |
| sampled descriptor rows only; it falls back internally (still exact) to the host extraction | |
| from the resident NMS / descriptor maps for other thresholds / borders or > 1024 candidates. | |
| Radii 1..8 other than the traced one: a precompiled per-radius device NMS + keypoint trace | |
| (taxonomy class B, built on first use). Radius 0 or > 8: host NMS on the traced scores.""" | |
| host_in = model.prepare_source(r8) # 480x640 plane, or a full-size plane for the device resize | |
| if model.kpc_ready and model.supports_device_nms_radius(nms_radius): | |
| # traced radius: one trace; other radii 1..8: + the precompiled per-radius NMS/keypoint | |
| # trace (captured on first use of that radius), replayed after the main trace | |
| with LOCK, torch.inference_mode(): | |
| kp, sc, desc = model.run_fused_keypoints_kpc( | |
| tt_in, host_in, keypoint_threshold=keypoint_threshold, max_keypoints=max_keypoints, | |
| border_removal_distance=border, with_descriptors=return_descriptors, nms_radius=nms_radius, | |
| ) | |
| return kp, sc, desc, True | |
| with LOCK, torch.inference_mode(): | |
| res = model.run_fused_prepared(tt_in, host_in, nms_radius=nms_radius) | |
| with torch.inference_mode(): | |
| if res.nms_map is not None: | |
| kp, sc, desc = _post.postprocess_from_nms_map( | |
| res.nms_map, res.descriptors_nchw, keypoint_threshold=keypoint_threshold, | |
| max_keypoints=max_keypoints, border_removal_distance=border, with_descriptors=return_descriptors, | |
| )[0] | |
| return kp, sc, desc, True | |
| kp, sc, desc = _post.postprocess_keypoints( | |
| res.scores_nchw, res.descriptors_nchw, nms_radius=nms_radius, keypoint_threshold=keypoint_threshold, | |
| max_keypoints=max_keypoints, border_removal_distance=border, with_descriptors=return_descriptors, | |
| )[0] | |
| return kp, sc, desc, False | |
| def _npz_b64(**arrays: np.ndarray) -> str: | |
| buf = io.BytesIO() | |
| np.savez(buf, **arrays) | |
| return base64.b64encode(buf.getvalue()).decode("ascii") | |
| # --------------------------------------------------------------------------- routes | |
| def health() -> dict: | |
| cfg = STATE.get("cfg") or {} | |
| return { | |
| "status": "ok" if STATE.get("ready") else "starting", | |
| "model": MODEL_NAME, | |
| "device": { | |
| "arch": "blackhole", | |
| "id": cfg.get("device_id", int(os.environ.get("TT_DEVICE_ID", "0"))), | |
| "open": STATE.get("device") is not None, | |
| }, | |
| } | |
| def info() -> dict: | |
| cfg = STATE.get("cfg") or {} | |
| mc = STATE.get("model_config") or {} | |
| return { | |
| "model": MODEL_NAME, | |
| "task": TASK, | |
| "io": "one image (base64 PNG/JPEG) -> keypoints [x, y] in original pixel coords, scores, optional 256-d descriptors", | |
| "hardware": "Tenstorrent Blackhole p150a, single chip (mesh 1x1) via tt-nn", | |
| "weights": { | |
| "repo": cfg.get("weights_repo", DEFAULT_WEIGHTS_REPO), | |
| "revision": cfg.get("weights_revision"), | |
| "local_dir": cfg.get("weights_dir"), | |
| "loaded": bool(STATE.get("ready")), | |
| }, | |
| "source": {"repo": SOURCE_REPO, "commit": SOURCE_COMMIT}, | |
| "device_config": { | |
| "dispatch": cfg.get("dispatch"), | |
| "num_command_queues": cfg.get("num_command_queues"), | |
| "compute_grid": cfg.get("compute_grid"), | |
| }, | |
| "input": { | |
| "canonical_height": INPUT_HEIGHT, | |
| "canonical_width": INPUT_WIDTH, | |
| "batch": 1, | |
| "preprocess": "bilinear resize to 640x480, /255, channel 0 (R) -- HF SuperPointImageProcessor defaults", | |
| }, | |
| "defaults": { | |
| "max_keypoints": 1024, | |
| "keypoint_threshold": 0.005, | |
| "nms_radius": 4, | |
| "border_removal_distance": mc.get("border_removal_distance", 4), | |
| "return_descriptors": True, | |
| }, | |
| "limits": {"max_keypoints": MAX_KEYPOINTS_CAP, "nms_radius": 32, "images_per_request": 1}, | |
| "serving_path": _serving_path(cfg), | |
| "warmup_ms": STATE.get("warmup_ms"), | |
| "descriptors": {"dim": 256, "encoding": "npz(base64) key 'descriptors' float16 (N, 256), L2-normalised"}, | |
| "license": LICENSE_NOTE, | |
| } | |
| def _serving_path(cfg: Dict[str, Any]) -> Dict[str, Any]: | |
| if not cfg.get("fused"): | |
| return { | |
| "traced": False, | |
| "device_nms": False, | |
| "custom_kernel": False, | |
| "device_softmax": True, | |
| "device_descriptor_l2norm": True, | |
| "note": "pure ttnn, untraced, host single-pass NMS (~6 fps device forward + ~36 ms host NMS); " | |
| "the README's 40.7 fps needs trace + the fused sp_eq_mul_mask kernel, not shipped here", | |
| } | |
| stages = list(cfg.get("fused_stages") or []) | |
| model = STATE.get("model") | |
| grid = None | |
| try: | |
| g = model.device.compute_with_storage_grid_size() if model is not None else None | |
| grid = f"{g.x}x{g.y}" if g is not None else None | |
| except Exception: # noqa: BLE001 - informational only | |
| grid = None | |
| return { | |
| "traced": True, | |
| "device_nms": "nms" in stages, | |
| # tt-nn generic_op kernels of this repo (code/kernels): block-0/1 cell convs + pools, merged head op | |
| # (score 1x1 + softmax inside), NMS fold / window max + keypoint candidates, descriptor sampler, | |
| # device bilinear resize | |
| "custom_kernel": True, | |
| "compute_grid": grid, | |
| "device_keypoints": bool(getattr(model, "kpc_ready", False)), | |
| "device_resize": bool(getattr(model, "device_resize", False)), | |
| "host_zero_copy_io": bool(getattr(model, "host_zc", False)), | |
| "device_softmax": True, | |
| "device_descriptor_l2norm": True, | |
| "fused": True, | |
| "fused_stages": stages, | |
| "nms_radius_traced": (getattr(model, "nms_radius_traced", 4) if "nms" in stages else None), | |
| "device_nms_radii": "traced radius in the main trace; 1..8 as precompiled per-radius traces " | |
| "(built on first use or at startup via SP_NMS_RADII_PRECOMPILE); 0 and > 8 on the host", | |
| "wide_page_upload": "wide" in stages, | |
| "rms_norm_l2": "rms" in stages, | |
| "row_major_outputs": "rm" in stages, | |
| "trace_region_size": cfg.get("trace_region_size"), | |
| "note": "TT_FUSED default: one metal trace per request (uint8 R plane in; full-size planes of the " | |
| "precompiled sizes are resized on device; encoder + heads + softmax, device NMS, keypoint list " | |
| "and bilinear descriptor sampling); one D2H of the keypoint header + sampled descriptor rows, " | |
| "host L2-normalise of those rows; nms_radius 1..8 other than the traced one replays a " | |
| "precompiled per-radius NMS trace, 0 and > 8 use the host NMS; TT_FUSED=0 restores the " | |
| "untraced host-NMS path", | |
| } | |
| def v1_models() -> dict: | |
| cfg = STATE.get("cfg") or {} | |
| return { | |
| "object": "list", | |
| "data": [{"id": cfg.get("weights_repo", DEFAULT_WEIGHTS_REPO), "object": "model", "owned_by": "changh95"}], | |
| } | |
| def _json_response(resp: Dict[str, Any]) -> Response: | |
| """The body FastAPI (>= 0.13x, ``-> dict`` route) renders for ``resp``: its fast path validates the | |
| dict against ``dict`` (identity for str / int / float / bool / list / dict) and serialises with | |
| pydantic-core's ``dump_json``; ``pydantic_core.to_json`` is that serialiser without the validation | |
| and threadpool hop, byte-identical (same float text, e.g. 9.3e-05 -> 0.000093).""" | |
| return Response(content=pydantic_core.to_json(resp), media_type="application/json") | |
| class PlaneParams(BaseModel): | |
| """Query parameters of the binary routes (same fields / limits as PredictRequest minus ``image``).""" | |
| max_keypoints: int = Field(1024, ge=-1, le=MAX_KEYPOINTS_CAP) | |
| keypoint_threshold: float = Field(0.005, ge=0.0, le=1.0) | |
| nms_radius: int = Field(4, ge=0, le=32) | |
| return_descriptors: bool = True | |
| def predict(req: PredictRequest) -> Response: | |
| return _json_response(predict_dict(req)) | |
| async def predict_raw(request: Request, max_keypoints: int = Query(1024, ge=-1, le=MAX_KEYPOINTS_CAP), | |
| keypoint_threshold: float = Query(0.005, ge=0.0, le=1.0), nms_radius: int = Query(4, ge=0, le=32), | |
| return_descriptors: bool = Query(True)) -> Response: | |
| """Body = the PNG/JPEG file bytes (application/octet-stream), parameters as query args; same | |
| response as /predict (no base64 / JSON request framing).""" | |
| raw = await request.body() | |
| p = PlaneParams(max_keypoints=max_keypoints, keypoint_threshold=keypoint_threshold, nms_radius=nms_radius, | |
| return_descriptors=return_descriptors) | |
| resp = await run_in_threadpool(_predict_core, p, lambda: _open_image(raw)) | |
| return _json_response(resp) | |
| async def predict_plane(request: Request, height: int = Query(..., ge=1, le=8192), width: int = Query(..., ge=1, le=8192), | |
| max_keypoints: int = Query(1024, ge=-1, le=MAX_KEYPOINTS_CAP), | |
| keypoint_threshold: float = Query(0.005, ge=0.0, le=1.0), nms_radius: int = Query(4, ge=0, le=32), | |
| return_descriptors: bool = Query(True)) -> Response: | |
| """Body = the image's channel 0 (R; the luma plane of a grayscale image) as raw uint8, row-major | |
| ``height x width`` -- the only channel the model reads. Skips the image decode; same response as | |
| /predict for the image whose R plane this is (the R channel of the image after PIL's ``convert("RGB")``; | |
| for L / LA / P-with-gray images that is the gray plane). Source sizes the device resize does not cover | |
| take the same host Pillow resize fallback as /predict (``TtSuperPoint.prepare_source``).""" | |
| n = height * width | |
| cl = request.headers.get("content-length") | |
| if cl is not None and cl.isdigit() and int(cl) != n: # reject before reading a wrong-size body | |
| raise HTTPException(status_code=400, detail=f"bad plane: {int(cl)} bytes, expected height*width={n}") | |
| raw = await request.body() | |
| if len(raw) != n: | |
| raise HTTPException(status_code=400, detail=f"bad plane: {len(raw)} bytes, expected height*width={height * width}") | |
| p = PlaneParams(max_keypoints=max_keypoints, keypoint_threshold=keypoint_threshold, nms_radius=nms_radius, | |
| return_descriptors=return_descriptors) | |
| plane = np.frombuffer(raw, dtype=np.uint8).reshape(height, width) | |
| resp = await run_in_threadpool(_predict_core, p, None, plane) | |
| return _json_response(resp) | |
| def predict_dict(req: PredictRequest) -> Dict[str, Any]: | |
| """/predict's response as a dict (the benches call this).""" | |
| return _predict_core(req, lambda: _decode_image(req.image)) | |
| def _predict_core(req, open_image, plane: np.ndarray | None = None) -> Dict[str, Any]: | |
| if not STATE.get("ready"): | |
| raise HTTPException(status_code=503, detail="model is still starting") | |
| model = STATE["model"] | |
| tt_in = STATE["tt_in"] | |
| border = STATE["model_config"]["border_removal_distance"] | |
| fused = bool(STATE.get("fused")) | |
| t0 = time.perf_counter() | |
| if plane is not None: | |
| orig_h, orig_w = plane.shape | |
| if fused: | |
| if model.device_resize: | |
| r8 = np.array(plane) # writable copy (prepare_source may hand it to torch) | |
| else: | |
| r8 = np.array(Image.fromarray(plane).resize((INPUT_WIDTH, INPUT_HEIGHT), resample=Image.BILINEAR), dtype=np.uint8) | |
| else: | |
| pixel_values = _preprocess(Image.fromarray(plane).convert("RGB")) | |
| else: | |
| im = open_image() | |
| orig_w, orig_h = im.size | |
| if fused: | |
| r8 = _r_plane(im) if model.device_resize else _preprocess_r8(im) | |
| else: | |
| pixel_values = _preprocess(im) | |
| t1 = time.perf_counter() | |
| device_nms = False | |
| try: | |
| if fused: | |
| # device_forward = H2D + trace + keypoint/descriptor readback (+ the host fallbacks) | |
| kp, sc, desc, device_nms = _infer_fused( | |
| model, tt_in, r8, max_keypoints=req.max_keypoints, keypoint_threshold=req.keypoint_threshold, | |
| nms_radius=req.nms_radius, return_descriptors=req.return_descriptors, border=border, | |
| ) | |
| t2 = time.perf_counter() | |
| else: | |
| with LOCK, torch.inference_mode(): | |
| scores_nchw, desc_nchw = model.run_untraced(tt_in, pixel_values) | |
| t2 = time.perf_counter() | |
| with torch.inference_mode(): | |
| kp, sc, desc = _post.postprocess_keypoints( | |
| scores_nchw, desc_nchw, | |
| nms_radius=req.nms_radius, | |
| keypoint_threshold=req.keypoint_threshold, | |
| max_keypoints=req.max_keypoints, | |
| border_removal_distance=border, | |
| with_descriptors=req.return_descriptors, | |
| )[0] | |
| except HTTPException: | |
| raise | |
| except Exception as e: | |
| LOG.exception("inference failed") | |
| raise HTTPException(status_code=500, detail=f"{type(e).__name__}: {e}") from None | |
| # Deterministic order: descending score (the port only sorts when top-k truncates). | |
| order = torch.argsort(sc, descending=True) | |
| kp, sc = kp[order], sc[order] | |
| if desc is not None: | |
| desc = desc[order] | |
| t3 = time.perf_counter() | |
| # Map from the 480x640 network frame back to the client's image. | |
| sx, sy = orig_w / INPUT_WIDTH, orig_h / INPUT_HEIGHT | |
| kp_np = kp.numpy().astype(np.float64) | |
| kp_orig = kp_np * np.array([sx, sy], dtype=np.float64) | |
| resp: Dict[str, Any] = { | |
| "num_keypoints": int(kp_np.shape[0]), | |
| "keypoints": [[round(x, 3), round(y, 3)] for x, y in kp_orig.tolist()], # tolist: same Python floats | |
| "scores": [round(s, 6) for s in sc.tolist()], # descending | |
| "original_size": {"height": orig_h, "width": orig_w}, | |
| "image_size": {"height": INPUT_HEIGHT, "width": INPUT_WIDTH}, | |
| "scale": {"x": sx, "y": sy}, | |
| "params": { | |
| "max_keypoints": req.max_keypoints, | |
| "keypoint_threshold": req.keypoint_threshold, | |
| "nms_radius": req.nms_radius, | |
| "border_removal_distance": border, | |
| }, | |
| **({"serving_path": {"traced": True, "device_nms": device_nms}} if fused else {}), | |
| "timing_ms": { | |
| "preprocess": round((t1 - t0) * 1000.0, 2), | |
| "device_forward": round((t2 - t1) * 1000.0, 2), | |
| "postprocess": round((t3 - t2) * 1000.0, 2), | |
| "total": round((t3 - t0) * 1000.0, 2), | |
| }, | |
| } | |
| if req.return_descriptors and desc is not None: | |
| desc16 = desc.to(torch.float16).numpy() # same RNE fp16 bits as numpy's astype, ~25x faster | |
| resp["descriptors"] = { | |
| "format": "npz", | |
| "key": "descriptors", | |
| "dtype": "float16", | |
| "shape": list(desc16.shape), | |
| "data": _npz_b64(descriptors=desc16), | |
| } | |
| return resp | |