Spaces:
Runtime error
Runtime error
| """ | |
| predict.py | |
| ========== | |
| End-to-end edge inference, fully decoupled from PyTorch: MediaPipe (landmark | |
| extraction) + ONNX Runtime (classification) only. | |
| Two entry points: | |
| * `predict_video(path)` -> list of {rank, gloss, confidence} (used by app.py) | |
| * CLI: python predict.py --video clip.mp4 | --hf data/data_0/00335.mp4 | |
| A module-level engine caches the MediaPipe model and the ONNX session so that | |
| repeated calls (e.g. a web backend) pay the ~13 MB model-load cost only once. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import shutil | |
| import tempfile | |
| import threading | |
| import time | |
| import urllib.request | |
| from pathlib import Path | |
| from typing import Dict, List, Optional, Tuple | |
| import numpy as np | |
| from slr.config import NUM_COORDS, NUM_LANDMARKS | |
| from slr.landmarks import ExtractorConfig, HolisticLandmarkExtractor | |
| RESOLVE = "https://huggingface.co/datasets/Voxel51/WLASL/resolve/main/{path}" | |
| # Defaults: the ONNX graph and the label map MUST correspond to the same trained | |
| # model. Override via env vars for a different deployment artifact. | |
| DEFAULT_ONNX = os.environ.get("SLR_ONNX", "onnx/sign_conformer_fp32.onnx") | |
| DEFAULT_LABELS = os.environ.get("SLR_LABELS", "data/landmarks_wlasl100/label_map.json") | |
| DEFAULT_WINDOW = int(os.environ.get("SLR_WINDOW", "96")) | |
| def _resolve_labels(path: Optional[str]) -> str: | |
| if path and Path(path).exists(): | |
| return path | |
| for cand in (DEFAULT_LABELS, | |
| "data/landmarks_canonical/label_map.json", | |
| "data/landmarks_wlasl100/label_map.json"): | |
| if Path(cand).exists(): | |
| return cand | |
| raise FileNotFoundError("No label_map.json found; pass labels_path explicitly.") | |
| def fit_window(clip: np.ndarray, window: int) -> Tuple[np.ndarray, np.ndarray]: | |
| """Pad/centre-crop a (T, V, C) clip to the window; return (x[1,W,V,C], mask[1,W]).""" | |
| T = clip.shape[0] | |
| if T >= window: | |
| s = (T - window) // 2 | |
| x = clip[s:s + window] | |
| mask = np.zeros(window, dtype=bool) | |
| else: | |
| pad = np.zeros((window - T, NUM_LANDMARKS, NUM_COORDS), np.float32) | |
| x = np.concatenate([clip, pad], 0) | |
| mask = np.zeros(window, dtype=bool); mask[T:] = True | |
| return x[None].astype(np.float32), mask[None] | |
| class InferenceEngine: | |
| """Holds the (reusable) MediaPipe extractor + ONNX session + label map.""" | |
| def __init__(self, onnx_path: str = DEFAULT_ONNX, | |
| labels_path: Optional[str] = None, | |
| window: int = DEFAULT_WINDOW): | |
| import onnxruntime as ort # local import keeps module import cheap | |
| self.window = window | |
| self.extractor = HolisticLandmarkExtractor(ExtractorConfig(normalize=True)) | |
| self.session = ort.InferenceSession( | |
| onnx_path, providers=["CPUExecutionProvider"]) | |
| label_map = json.load(open(_resolve_labels(labels_path), encoding="utf-8")) | |
| self.idx_to_gloss = {v: k for k, v in label_map.items()} | |
| self._lock = threading.Lock() # MediaPipe graph is not thread-safe | |
| def predict(self, video_path: str | Path, topk: int = 5) -> Dict: | |
| """Return {predictions:[{rank,gloss,confidence}], frames, timings}.""" | |
| t0 = time.time() | |
| with self._lock: | |
| clip, _ = self.extractor.extract(video_path) # (T, V, C) | |
| t_lm = (time.time() - t0) * 1000 | |
| x, mask = fit_window(clip, self.window) | |
| t1 = time.time() | |
| logits = self.session.run( | |
| ["logits"], {"landmarks": x, "pad_mask": mask})[0][0] | |
| t_cls = (time.time() - t1) * 1000 | |
| probs = np.exp(logits - logits.max()); probs /= probs.sum() | |
| top = probs.argsort()[::-1][:topk] | |
| preds = [{"rank": r + 1, | |
| "gloss": self.idx_to_gloss.get(int(i), "?"), | |
| "confidence": round(float(probs[i]) * 100, 2)} | |
| for r, i in enumerate(top)] | |
| return { | |
| "predictions": preds, | |
| "frames": int(clip.shape[0]), | |
| "timings_ms": {"landmark": round(t_lm, 1), "inference": round(t_cls, 1)}, | |
| } | |
| # --- module-level singleton (lazy) ----------------------------------------- | |
| _ENGINE: Optional[InferenceEngine] = None | |
| _ENGINE_LOCK = threading.Lock() | |
| def get_engine() -> InferenceEngine: | |
| global _ENGINE | |
| if _ENGINE is None: | |
| with _ENGINE_LOCK: | |
| if _ENGINE is None: | |
| _ENGINE = InferenceEngine() | |
| return _ENGINE | |
| def predict_video(video_path: str | Path, topk: int = 5) -> Dict: | |
| """Public API used by app.py: video file -> Top-k prediction dict.""" | |
| return get_engine().predict(video_path, topk=topk) | |
| # --- CLI ------------------------------------------------------------------- | |
| def _cli(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--video", help="local video file") | |
| ap.add_argument("--hf", help="repo-relative clip path to stream from the mirror") | |
| ap.add_argument("--topk", type=int, default=5) | |
| args = ap.parse_args() | |
| tmp = None | |
| if args.hf: | |
| tmp = tempfile.TemporaryDirectory() | |
| dst = Path(tmp.name) / "clip.mp4" | |
| url = RESOLVE.format(path=args.hf) | |
| with urllib.request.urlopen( | |
| urllib.request.Request(url, headers={"User-Agent": "slr/1.0"}), | |
| timeout=60) as r, open(dst, "wb") as f: | |
| shutil.copyfileobj(r, f) | |
| video = dst | |
| else: | |
| video = Path(args.video) | |
| res = predict_video(video, topk=args.topk) | |
| print(f"\nclip: {args.hf or args.video} ({res['frames']} frames)") | |
| print(f"landmark: {res['timings_ms']['landmark']} ms | " | |
| f"onnx: {res['timings_ms']['inference']} ms\n") | |
| print("Top predictions:") | |
| for p in res["predictions"]: | |
| bar = "#" * int(p["confidence"] / 100 * 30) | |
| print(f" {p['rank']}. {p['gloss']:<14} {p['confidence']:5.1f}% {bar}") | |
| if tmp: | |
| tmp.cleanup() | |
| if __name__ == "__main__": | |
| _cli() | |