File size: 38,286 Bytes
dc14621
3461c19
dc14621
 
 
 
 
 
3461c19
 
 
 
dc14621
0887b31
 
 
 
 
 
c699c4c
 
 
 
 
0887b31
 
 
 
 
 
 
3461c19
 
 
 
 
 
 
 
 
026da6c
 
 
 
0887b31
c699c4c
 
 
0887b31
c699c4c
 
3461c19
 
 
dc14621
 
 
 
c699c4c
dc14621
3461c19
dc14621
3461c19
6ffd3f8
3461c19
 
dc14621
3461c19
dc14621
 
6ffd3f8
dc14621
6ffd3f8
 
3461c19
6ffd3f8
dc14621
 
 
3461c19
026da6c
6ffd3f8
 
dc14621
3461c19
dc14621
2237e2c
3461c19
 
 
 
 
 
 
 
 
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
 
dc14621
 
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
dc14621
3461c19
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
0887b31
3461c19
 
 
 
 
 
0887b31
 
 
026da6c
3461c19
 
dc14621
3461c19
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dc14621
 
 
 
3461c19
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6ffd3f8
3461c19
0887b31
 
 
026da6c
 
 
 
0887b31
026da6c
0887b31
026da6c
0887b31
 
 
026da6c
 
 
 
 
 
 
 
 
 
 
 
dc14621
3461c19
dc14621
0887b31
 
 
3461c19
 
 
0887b31
3461c19
 
 
0887b31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3461c19
dc14621
 
3461c19
 
 
dc14621
0887b31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0887b31
 
 
 
 
 
 
 
 
3461c19
 
 
 
 
0887b31
3461c19
 
 
 
0887b31
 
3461c19
 
 
 
 
 
 
dc14621
 
 
 
 
3461c19
 
 
 
 
 
 
 
 
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
dc14621
3461c19
dc14621
 
c699c4c
 
 
6ffd3f8
 
 
c699c4c
 
 
 
 
 
6ffd3f8
 
 
 
 
 
c699c4c
 
 
 
 
 
3461c19
dc14621
c699c4c
6ffd3f8
 
 
 
 
 
 
3461c19
 
c699c4c
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3461c19
 
 
 
 
 
 
dc14621
 
 
 
3461c19
 
 
 
 
 
 
 
 
 
dc14621
 
 
 
3461c19
 
dc14621
3461c19
 
 
 
 
 
 
 
 
 
 
026da6c
 
 
 
 
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
0887b31
 
 
 
 
 
 
 
 
 
3461c19
 
 
 
 
 
 
0887b31
 
 
c699c4c
 
 
 
 
 
0887b31
 
 
c699c4c
 
 
 
 
 
 
 
0887b31
 
 
 
 
c699c4c
 
0887b31
 
 
 
c699c4c
 
 
 
 
 
3461c19
 
 
 
 
 
 
 
 
dc14621
 
 
6ffd3f8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dc14621
6ffd3f8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3461c19
dc14621
3461c19
 
 
dc14621
c699c4c
3461c19
6ffd3f8
 
 
 
 
 
 
 
 
c699c4c
6ffd3f8
 
 
 
 
 
3461c19
dc14621
0887b31
3461c19
0887b31
c699c4c
 
 
 
 
0887b31
 
 
 
 
 
 
 
 
 
 
 
 
 
3461c19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6ffd3f8
 
3461c19
 
 
 
 
 
 
 
 
0887b31
3461c19
 
 
 
 
 
 
 
6ffd3f8
3461c19
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
# 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


@asynccontextmanager
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)


@app.exception_handler(RequestValidationError)
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


@app.get("/health")
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,
        },
    }


@app.get("/info")
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",
    }


@app.get("/v1/models")
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


@app.post("/predict")
def predict(req: PredictRequest) -> Response:
    return _json_response(predict_dict(req))


@app.post("/predict_raw")
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)


@app.post("/predict_plane")
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