Download code/models/server/smoke_test.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 9.34 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/server/smoke_test.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/server/smoke_test.py
-
curl -L -o smoke_test.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/server/smoke_test.py
9.34 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Smoke test for the SuperPoint tt-dit-server: one request, sanity checks, PASS/FAIL. | |
| python code/models/server/smoke_test.py --url http://127.0.0.1:20000 | |
| Default input: ``sample_data/house_in_field_1080p.jpg`` (the port's natural-image | |
| validation frame, 1600x900) resolved relative to this file (code/sample_data). | |
| Only the standard library + numpy are needed (numpy decodes the descriptor NPZ). | |
| Checks (the README reports ~500 strong keypoints on this frame at top-500; | |
| we ask for 1024 and require a few hundred): | |
| * /health is ok, /info answers, /v1/models lists the weights repo | |
| * /predict returns 200 with num_keypoints in [200, max_keypoints] | |
| * every keypoint lies inside the ORIGINAL image, scores in (threshold, 1] | |
| and sorted-consistent with top-k, descriptors (N, 256) float16 finite and | |
| unit-norm (|1 - ||d||| < 0.05), timing_ms present | |
| * a second call with return_descriptors=false / max_keypoints=100 obeys both | |
| * POST /predict_raw (the file bytes) and POST /predict_plane (the R plane, needs Pillow) return the | |
| same response as /predict apart from timing_ms (skip with --no-binary-routes) | |
| Exit code 0 on PASS, 1 on FAIL. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import base64 | |
| import io | |
| import json | |
| import sys | |
| import time | |
| import urllib.error | |
| import urllib.request | |
| from pathlib import Path | |
| DEFAULT_IMAGE = Path(__file__).resolve().parents[2] / "sample_data" / "house_in_field_1080p.jpg" | |
| def _get(url: str, timeout: float = 30.0): | |
| with urllib.request.urlopen(url, timeout=timeout) as r: | |
| return r.status, json.loads(r.read().decode("utf-8")) | |
| def _post(url: str, payload: dict, timeout: float = 300.0): | |
| data = json.dumps(payload).encode("utf-8") | |
| req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"}, method="POST") | |
| try: | |
| with urllib.request.urlopen(req, timeout=timeout) as r: | |
| return r.status, json.loads(r.read().decode("utf-8")) | |
| except urllib.error.HTTPError as e: | |
| body = e.read().decode("utf-8", "replace") | |
| return e.code, {"detail": body} | |
| def _post_bytes(url: str, body: bytes, timeout: float = 300.0): | |
| req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/octet-stream"}, method="POST") | |
| try: | |
| with urllib.request.urlopen(req, timeout=timeout) as r: | |
| return r.status, json.loads(r.read().decode("utf-8")) | |
| except urllib.error.HTTPError as e: | |
| return e.code, {"detail": e.read().decode("utf-8", "replace")} | |
| def _fail(msg: str) -> int: | |
| print(f"FAIL {msg}") | |
| return 1 | |
| def main() -> int: | |
| ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) | |
| ap.add_argument("--url", default="http://127.0.0.1:20000", help="server base URL (the port `serve` printed)") | |
| ap.add_argument("--image", type=Path, default=DEFAULT_IMAGE) | |
| ap.add_argument("--max-keypoints", type=int, default=1024) | |
| ap.add_argument("--min-keypoints", type=int, default=200, help="lower bound for PASS on the default image") | |
| ap.add_argument("--threshold", type=float, default=0.005) | |
| ap.add_argument("--timeout", type=float, default=300.0) | |
| ap.add_argument("--no-binary-routes", action="store_true", help="skip /predict_raw and /predict_plane") | |
| ap.add_argument("--save-json", type=Path, default=None, help="write the first /predict response here") | |
| args = ap.parse_args() | |
| base = args.url.rstrip("/") | |
| if not args.image.is_file(): | |
| return _fail(f"input image not found: {args.image}") | |
| raw = args.image.read_bytes() | |
| if raw[:4] not in (b"\xff\xd8\xff\xe0", b"\xff\xd8\xff\xe1", b"\x89PNG") and raw[:3] != b"\xff\xd8\xff": | |
| return _fail(f"{args.image} is not a PNG/JPEG (git-lfs pointer? run `git lfs pull`)") | |
| b64 = base64.b64encode(raw).decode("ascii") | |
| try: | |
| st, health = _get(f"{base}/health") | |
| except Exception as e: | |
| return _fail(f"GET /health failed: {e}") | |
| if st != 200 or health.get("status") != "ok": | |
| return _fail(f"/health not ok: {health}") | |
| st, info = _get(f"{base}/info") | |
| if st != 200 or info.get("model") != "superpoint-p150": | |
| return _fail(f"/info unexpected: {info}") | |
| st, models = _get(f"{base}/v1/models") | |
| if st != 200 or not models.get("data"): | |
| return _fail(f"/v1/models unexpected: {models}") | |
| weights_repo = models["data"][0]["id"] | |
| payload = { | |
| "image": b64, | |
| "max_keypoints": args.max_keypoints, | |
| "keypoint_threshold": args.threshold, | |
| "nms_radius": 4, | |
| "return_descriptors": True, | |
| } | |
| t0 = time.perf_counter() | |
| st, resp = _post(f"{base}/predict", payload, timeout=args.timeout) | |
| wall_ms = (time.perf_counter() - t0) * 1000.0 | |
| if st != 200: | |
| return _fail(f"/predict HTTP {st}: {str(resp)[:300]}") | |
| if args.save_json: | |
| args.save_json.write_text(json.dumps(resp)) | |
| n = resp.get("num_keypoints") | |
| kps = resp.get("keypoints") or [] | |
| scores = resp.get("scores") or [] | |
| if not isinstance(n, int) or n != len(kps) or n != len(scores): | |
| return _fail(f"inconsistent counts: num_keypoints={n} keypoints={len(kps)} scores={len(scores)}") | |
| if n < args.min_keypoints or n > args.max_keypoints: | |
| return _fail(f"num_keypoints={n} outside [{args.min_keypoints}, {args.max_keypoints}]") | |
| osz = resp.get("original_size") or {} | |
| W, H = osz.get("width"), osz.get("height") | |
| if not W or not H: | |
| return _fail(f"original_size missing: {osz}") | |
| for x, y in kps: | |
| if not (0.0 <= x < W and 0.0 <= y < H): | |
| return _fail(f"keypoint ({x}, {y}) outside the {W}x{H} image") | |
| if min(scores) <= args.threshold or max(scores) > 1.0 + 1e-6: | |
| return _fail(f"scores out of range: min={min(scores)} max={max(scores)} threshold={args.threshold}") | |
| if "timing_ms" not in resp or "device_forward" not in resp["timing_ms"]: | |
| return _fail("timing_ms.device_forward missing") | |
| d = resp.get("descriptors") | |
| if not d or d.get("format") != "npz" or d.get("shape") != [n, 256]: | |
| return _fail(f"descriptors missing or wrong shape: {None if not d else d.get('shape')}") | |
| try: | |
| import numpy as np | |
| npz = np.load(io.BytesIO(base64.b64decode(d["data"]))) | |
| desc = npz[d.get("key", "descriptors")] | |
| if desc.shape != (n, 256) or desc.dtype != np.float16: | |
| return _fail(f"descriptor array is {desc.shape} {desc.dtype}, expected ({n}, 256) float16") | |
| desc32 = desc.astype(np.float32) | |
| if not np.isfinite(desc32).all(): | |
| return _fail("descriptors contain non-finite values") | |
| norms = np.linalg.norm(desc32, axis=1) | |
| if np.abs(1.0 - norms).max() > 0.05: | |
| return _fail(f"descriptors not unit-norm: |1-norm| max {np.abs(1.0 - norms).max():.3f}") | |
| norm_dev = float(np.abs(1.0 - norms).max()) | |
| except ImportError: | |
| norm_dev = float("nan") # numpy unavailable: shape/format checks above still ran | |
| # Second call: descriptors off + small top-k must be honoured. | |
| st2, resp2 = _post(f"{base}/predict", {**payload, "return_descriptors": False, "max_keypoints": 100}, timeout=args.timeout) | |
| if st2 != 200: | |
| return _fail(f"second /predict HTTP {st2}: {str(resp2)[:200]}") | |
| if resp2.get("num_keypoints") != 100 or "descriptors" in resp2: | |
| return _fail(f"top-k/return_descriptors not honoured: n={resp2.get('num_keypoints')} desc={'descriptors' in resp2}") | |
| binary = "skipped" | |
| if not args.no_binary_routes: | |
| from urllib.parse import urlencode | |
| q = {k: v for k, v in payload.items() if k != "image"} | |
| q["return_descriptors"] = "true" | |
| ref = {k: v for k, v in resp.items() if k != "timing_ms"} | |
| st3, resp3 = _post_bytes(f"{base}/predict_raw?{urlencode(q)}", raw, timeout=args.timeout) | |
| if st3 != 200 or {k: v for k, v in resp3.items() if k != "timing_ms"} != ref: | |
| return _fail(f"/predict_raw HTTP {st3}, response differs from /predict: {str(resp3)[:200]}") | |
| binary = "raw=ok" | |
| try: | |
| from PIL import Image | |
| except ImportError: | |
| binary += " plane=skipped(no Pillow)" | |
| else: | |
| im = Image.open(io.BytesIO(raw)) | |
| im = im if im.mode == "RGB" else im.convert("RGB") | |
| plane = im.getchannel(0).tobytes() | |
| st4, resp4 = _post_bytes(f"{base}/predict_plane?{urlencode({**q, 'height': im.height, 'width': im.width})}", | |
| plane, timeout=args.timeout) | |
| if st4 != 200 or {k: v for k, v in resp4.items() if k != "timing_ms"} != ref: | |
| return _fail(f"/predict_plane HTTP {st4}, response differs from /predict: {str(resp4)[:200]}") | |
| binary += " plane=ok" | |
| tm = resp["timing_ms"] | |
| print( | |
| f"PASS superpoint-p150: {n} keypoints on {args.image.name} ({W}x{H}), " | |
| f"score max={max(scores):.4f} min={min(scores):.4f}, desc (N,256) f16 |1-norm|max={norm_dev:.3f}, " | |
| f"device_forward={tm.get('device_forward')} ms postprocess={tm.get('postprocess')} ms " | |
| f"server_total={tm.get('total')} ms wall={wall_ms:.0f} ms, binary routes {binary}, weights={weights_repo}" | |
| ) | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |