Add Tenstorrent Blackhole tt-nn port
Browse files- SERVING.md +80 -0
- code/models/server/__init__.py +0 -0
- code/models/server/app.py +137 -0
- tt-model.yaml +73 -0
SERVING.md
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Serving this port with tt-model-manager
|
| 2 |
+
|
| 3 |
+
This repo carries everything needed to build a **tt-model container bundle** from it, so
|
| 4 |
+
`tt-model pull` / `tt-model serve` and `tt-cli` can run it. Do this on a
|
| 5 |
+
**Blackhole host** (amd64 Linux + Docker >= 25).
|
| 6 |
+
|
| 7 |
+
## Why not on a Mac
|
| 8 |
+
|
| 9 |
+
`tt-model package --container` builds from `ghcr.io/.../ubuntu-22.04-dev-amd64` and a cold
|
| 10 |
+
build is 2.5-4 hours. Packaging and pushing need **no card** (`require_host(need_devices=False)`),
|
| 11 |
+
but they do need amd64. `tt-model serve` needs `/dev/tenstorrent/*`.
|
| 12 |
+
|
| 13 |
+
## Steps
|
| 14 |
+
|
| 15 |
+
```bash
|
| 16 |
+
# 0. get tt-model (not on PyPI)
|
| 17 |
+
git clone https://github.com/tenstorrent/tt-model-manager && cd tt-model-manager
|
| 18 |
+
python -m venv .venv && source .venv/bin/activate && pip install -e .
|
| 19 |
+
|
| 20 |
+
# 1. pull THIS repo
|
| 21 |
+
hf download changh95/superpoint-blackhole --local-dir superpoint-blackhole && cd superpoint-blackhole
|
| 22 |
+
|
| 23 |
+
# 2. point the manifest at your tt-metal checkout
|
| 24 |
+
$EDITOR tt-model.yaml # source.tt_metal: /path/to/tt-metal
|
| 25 |
+
|
| 26 |
+
# 3. finish the serving adapter
|
| 27 |
+
$EDITOR code/models/server/app.py # see "VERIFY ON HARDWARE" markers
|
| 28 |
+
|
| 29 |
+
# 4. validate with no hardware and no build
|
| 30 |
+
python -c "from tt_kernel.container_manifest import ContainerManifest as M; M.load('tt-model.yaml')"
|
| 31 |
+
|
| 32 |
+
# 5. build, prove, publish
|
| 33 |
+
tt-model package --container tt-model.yaml
|
| 34 |
+
tt-model serve build/superpoint-blackhole/tt_kernel_manifest.json
|
| 35 |
+
tt-model push build/superpoint-blackhole --publish
|
| 36 |
+
```
|
| 37 |
+
|
| 38 |
+
Step 5's `--publish` implies `--public` and adds the `tt-model-catalog` tag, which is what
|
| 39 |
+
lists it in the community catalog.
|
| 40 |
+
|
| 41 |
+
## What the bundle push does to this repo
|
| 42 |
+
|
| 43 |
+
- **`code/` and `image/` are replaced wholesale** (`_prune_removed` in `tt_kernel/hub.py`).
|
| 44 |
+
The port code staged here is exactly what `extra_code` re-stages, so it is replaced, not
|
| 45 |
+
duplicated. Files under `code/` that the allowlist does not ship get pruned.
|
| 46 |
+
- **`README.md` is overwritten** by the generated card. The text you want to keep belongs in
|
| 47 |
+
`card.description` / `card.quickstart` in `tt-model.yaml` -- it is already seeded there.
|
| 48 |
+
- **Root files survive**, which is why `media/` and this file live at the root.
|
| 49 |
+
|
| 50 |
+
## The serving contract
|
| 51 |
+
|
| 52 |
+
`kind: tt-dit-server` requires only `hardware` + `mesh_device` (no `max_num_seqs` /
|
| 53 |
+
`block_size` -- those are vLLM engine settings). `runtime.app` is an ASGI target
|
| 54 |
+
`"module:attr"` whose **top-level package must appear in an allowlist entry**, which is why
|
| 55 |
+
`runtime.app` and `source.extra_code[0].paths` have to stay in sync. The kind installs
|
| 56 |
+
fastapi / uvicorn / pydantic / pillow; `runtime.packages` adds to that set.
|
| 57 |
+
|
| 58 |
+
Despite the name, this kind is not diffusion-only -- `jashansinghTT/s2-pro-blackhole` is a
|
| 59 |
+
TTS model using it. It simply means "the model's own ASGI app under uvicorn".
|
| 60 |
+
|
| 61 |
+
## Known sharp edges (offline-validated, not yet hardware-validated)
|
| 62 |
+
|
| 63 |
+
The manifest loads clean (`load_container_manifest(..., check_sources=False)` passes, which
|
| 64 |
+
includes the launcher's own `validate()`), but two things can only be settled on the box:
|
| 65 |
+
|
| 66 |
+
1. **`code/` becomes the image's ONLY `models` package.** `source.code: [models/common]` is
|
| 67 |
+
staged from tt-metal into `code/models/common`, and an `extra_code` path of `models`
|
| 68 |
+
merges the port's own tree into the same `code/models/`. tt-metal's `models/` has no
|
| 69 |
+
`__init__.py`, so this relies on namespace packages resolving. Check the import in
|
| 70 |
+
`verify.sh` output on the first build.
|
| 71 |
+
2. **Generically-named top-level packages.** Repos whose package is `models`, `tt`, or
|
| 72 |
+
`common` put a very common name at the root of the image's import path. If an import
|
| 73 |
+
collides, rename the package in the port repo (and update `runtime.app` +
|
| 74 |
+
`source.extra_code[].paths` together -- the launcher cross-checks them).
|
| 75 |
+
|
| 76 |
+
`source.code` needs at least one tt-metal-relative entry even when all real code comes from
|
| 77 |
+
`extra_code`; `models/common` is used as that minimum. If your port genuinely imports more
|
| 78 |
+
from tt-metal (`locate-anything` needs `models/tt_transformers` and
|
| 79 |
+
`models/demos/qwen25_vl`), list it there instead -- under-listing fails the image's own
|
| 80 |
+
build-time import check on your machine, which is the cheap place to find out.
|
code/models/server/__init__.py
ADDED
|
File without changes
|
code/models/server/app.py
ADDED
|
@@ -0,0 +1,137 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""ASGI serving app for SuperPoint on Tenstorrent Blackhole.
|
| 3 |
+
|
| 4 |
+
Served by tt-model-manager as ``kind: tt-dit-server``::
|
| 5 |
+
|
| 6 |
+
runtime:
|
| 7 |
+
app: models.server.app:app
|
| 8 |
+
|
| 9 |
+
uvicorn runs this module. The device is opened and the model built inside the ASGI
|
| 10 |
+
**lifespan**, so uvicorn's ``Application startup complete`` -- which is the line
|
| 11 |
+
``tt-model serve`` waits for -- means the chip is claimed and the model is warm.
|
| 12 |
+
|
| 13 |
+
I/O: image -> keypoints, scores, descriptors
|
| 14 |
+
"""
|
| 15 |
+
from __future__ import annotations
|
| 16 |
+
|
| 17 |
+
import base64
|
| 18 |
+
import io
|
| 19 |
+
import os
|
| 20 |
+
from contextlib import asynccontextmanager
|
| 21 |
+
from typing import Any
|
| 22 |
+
|
| 23 |
+
import numpy as np
|
| 24 |
+
import torch
|
| 25 |
+
from fastapi import FastAPI, HTTPException
|
| 26 |
+
from PIL import Image
|
| 27 |
+
from pydantic import BaseModel, Field
|
| 28 |
+
|
| 29 |
+
from models.tt.superpoint_ttnn import TtSuperPoint
|
| 30 |
+
|
| 31 |
+
STATE: dict[str, Any] = {}
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def _open_device():
|
| 35 |
+
"""Open the chip.
|
| 36 |
+
|
| 37 |
+
kwargs are this repo's own (from its conftest/demo), not generic defaults.
|
| 38 |
+
``TT_MESH_SHAPE`` is set by tt-model from the manifest's ``mesh_device``; these
|
| 39 |
+
ports are single-chip, so a shape other than (1, 1) needs open_mesh_device.
|
| 40 |
+
"""
|
| 41 |
+
import ttnn
|
| 42 |
+
|
| 43 |
+
shape = os.environ.get("TT_MESH_SHAPE", "(1, 1)")
|
| 44 |
+
if shape.replace(" ", "") not in ("(1,1)", "1,1"):
|
| 45 |
+
raise RuntimeError(
|
| 46 |
+
f"{shape} is a multi-chip mesh; this app opens a single device. "
|
| 47 |
+
"Use ttnn.open_mesh_device and shard the model first."
|
| 48 |
+
)
|
| 49 |
+
return ttnn.open_device(
|
| 50 |
+
device_id=int(os.environ.get("TT_DEVICE_ID", "0")),
|
| 51 |
+
l1_small_size=32768, trace_region_size=1500000000,
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _load_reference():
|
| 56 |
+
"""Load the torch reference this port wraps."""
|
| 57 |
+
from transformers import SuperPointForKeypointDetection
|
| 58 |
+
return SuperPointForKeypointDetection.from_pretrained('magic-leap-community/superpoint').eval()
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
@asynccontextmanager
|
| 62 |
+
async def lifespan(_app: FastAPI):
|
| 63 |
+
torch.set_grad_enabled(False)
|
| 64 |
+
device = _open_device()
|
| 65 |
+
STATE["device"] = device
|
| 66 |
+
try:
|
| 67 |
+
STATE["model"] = TtSuperPoint(torch_model=_load_reference(), device=device, input_height=480, input_width=640)
|
| 68 |
+
yield
|
| 69 |
+
finally:
|
| 70 |
+
import ttnn
|
| 71 |
+
|
| 72 |
+
STATE.pop("model", None)
|
| 73 |
+
ttnn.close_device(STATE.pop("device"))
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
app = FastAPI(title="SuperPoint on Blackhole", lifespan=lifespan)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class PredictRequest(BaseModel):
|
| 80 |
+
"""One inference request. ``image`` is base64-encoded (PNG/JPEG)."""
|
| 81 |
+
|
| 82 |
+
image: str
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def _decode(b64: str) -> torch.Tensor:
|
| 86 |
+
try:
|
| 87 |
+
raw = base64.b64decode(b64, validate=True)
|
| 88 |
+
im = Image.open(io.BytesIO(raw)).convert("RGB")
|
| 89 |
+
except Exception as e:
|
| 90 |
+
raise HTTPException(status_code=400, detail=f"bad image: {e}") from None
|
| 91 |
+
arr = np.asarray(im, dtype=np.float32) / 255.0
|
| 92 |
+
return torch.from_numpy(arr).permute(2, 0, 1)[None].float()
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
@app.get("/health")
|
| 96 |
+
def health() -> dict:
|
| 97 |
+
return {"status": "ok" if "model" in STATE else "starting"}
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
@app.get("/info")
|
| 101 |
+
def info() -> dict:
|
| 102 |
+
return {
|
| 103 |
+
"model": "SuperPoint",
|
| 104 |
+
"io": "image -> keypoints, scores, descriptors",
|
| 105 |
+
"hardware": "Tenstorrent Blackhole p150a (single chip)",
|
| 106 |
+
"weights": "magic-leap-community/superpoint",
|
| 107 |
+
"source": "https://github.com/changh95/tt-superpoint",
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
@app.post("/predict")
|
| 112 |
+
def predict(req: PredictRequest) -> dict:
|
| 113 |
+
model = STATE.get("model")
|
| 114 |
+
if model is None:
|
| 115 |
+
raise HTTPException(status_code=503, detail="model is still starting")
|
| 116 |
+
pixel_values = _decode(req.image)
|
| 117 |
+
out = model.forward(pixel_values)
|
| 118 |
+
return {"output": _jsonable(out)}
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def _jsonable(obj: Any) -> Any:
|
| 122 |
+
"""Tensors/arrays -> nested lists, so the response is plain JSON.
|
| 123 |
+
|
| 124 |
+
VERIFY ON HARDWARE: for large dense outputs (depth maps, pointmaps) returning
|
| 125 |
+
raw lists is wasteful -- prefer a PNG/NPZ response for those.
|
| 126 |
+
"""
|
| 127 |
+
if isinstance(obj, torch.Tensor):
|
| 128 |
+
return obj.detach().cpu().tolist()
|
| 129 |
+
if isinstance(obj, np.ndarray):
|
| 130 |
+
return obj.tolist()
|
| 131 |
+
if isinstance(obj, dict):
|
| 132 |
+
return {k: _jsonable(v) for k, v in obj.items()}
|
| 133 |
+
if isinstance(obj, (list, tuple)):
|
| 134 |
+
return [_jsonable(v) for v in obj]
|
| 135 |
+
if hasattr(obj, "_asdict"):
|
| 136 |
+
return _jsonable(obj._asdict())
|
| 137 |
+
return obj
|
tt-model.yaml
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# tt-model-manager container manifest (schema 5.1) for SuperPoint on Blackhole.
|
| 3 |
+
#
|
| 4 |
+
# tt-model package --container tt-model.yaml # amd64 Linux + Docker>=25, ~2.5-4 h
|
| 5 |
+
# tt-model serve build/superpoint-blackhole/tt_kernel_manifest.json # needs a card
|
| 6 |
+
# tt-model push build/superpoint-blackhole --publish # adds the tt-model-catalog tag
|
| 7 |
+
#
|
| 8 |
+
# SET ME before building: source.tt_metal must point at YOUR built tt-metal tree.
|
| 9 |
+
schema: "5.1"
|
| 10 |
+
|
| 11 |
+
repo: changh95/superpoint-blackhole
|
| 12 |
+
name: superpoint-blackhole
|
| 13 |
+
|
| 14 |
+
# A POINTER. Weights are never baked into the image; they land in the consumer's
|
| 15 |
+
# own HF cache under their own token. Pin a revision once validated.
|
| 16 |
+
weights: magic-leap-community/superpoint
|
| 17 |
+
|
| 18 |
+
kind: tt-dit-server
|
| 19 |
+
arch: blackhole
|
| 20 |
+
|
| 21 |
+
source:
|
| 22 |
+
# SET ME -- a local checkout is hermetic and lets tt-model read tt-metal's torch
|
| 23 |
+
# pin (metal_torch_pin), which verify.sh then asserts. A {repo, ref} mapping
|
| 24 |
+
# also works but skips that pin.
|
| 25 |
+
tt_metal: /path/to/tt-metal
|
| 26 |
+
|
| 27 |
+
# Relative to the tt-metal tree. At least one entry is required by the schema;
|
| 28 |
+
# this port's own code comes from extra_code below.
|
| 29 |
+
code:
|
| 30 |
+
- models/common
|
| 31 |
+
|
| 32 |
+
# This repo's code, which does NOT live in the tt-metal tree. `root: code` works
|
| 33 |
+
# because the published HF repo carries the port under code/ -- so on the box you
|
| 34 |
+
# just pull this repo and build from it.
|
| 35 |
+
extra_code:
|
| 36 |
+
- root: code
|
| 37 |
+
paths:
|
| 38 |
+
- models
|
| 39 |
+
|
| 40 |
+
ubuntu: "22.04"
|
| 41 |
+
python: "3.12"
|
| 42 |
+
|
| 43 |
+
runtime:
|
| 44 |
+
app: models.server.app:app
|
| 45 |
+
mesh_shape_env: TT_MESH_SHAPE
|
| 46 |
+
# Added ON TOP of the kind's defaults (fastapi, uvicorn, pydantic>=2, pillow).
|
| 47 |
+
# torch is auto-pinned to tt-metal's pin when tt_metal is a local path.
|
| 48 |
+
packages:
|
| 49 |
+
- numpy
|
| 50 |
+
- transformers
|
| 51 |
+
|
| 52 |
+
serve:
|
| 53 |
+
port: 20000
|
| 54 |
+
hardware: p150
|
| 55 |
+
mesh_device: P150
|
| 56 |
+
|
| 57 |
+
# Build-time assertions, run INSIDE the finished image. No device is opened.
|
| 58 |
+
verify:
|
| 59 |
+
- "import models.server.app as a; assert a.app"
|
| 60 |
+
|
| 61 |
+
card:
|
| 62 |
+
description: >
|
| 63 |
+
SuperPoint on a single Tenstorrent Blackhole p150a via tt-nn.
|
| 64 |
+
image -> keypoints, scores, descriptors.
|
| 65 |
+
quickstart: |
|
| 66 |
+
### Call it
|
| 67 |
+
|
| 68 |
+
```bash
|
| 69 |
+
curl -s localhost:20000/health
|
| 70 |
+
curl -s localhost:20000/info
|
| 71 |
+
curl -s localhost:20000/predict -H 'Content-Type: application/json' \
|
| 72 |
+
-d "{\"image\": \"$(base64 -w0 your.png)\"}"
|
| 73 |
+
```
|