changh95 commited on
Commit
dc14621
·
verified ·
1 Parent(s): f64e572

Add Tenstorrent Blackhole tt-nn port

Browse files
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
+ ```