Files
scene-actor-extraction/scripts/sae_embed_loader.py
T
dtourolleandClaude Opus 5 458116f118 fix(trt): drop explicit shapes for static TransNetV2; gallery over-fetch + dedup
trtexec rejects --minShapes/--optShapes/--maxShapes for a fully static model
("Static model does not take explicit shapes"). TransNetV2's input is fixed at
1x100x27x48x3, so the shape comes from the model itself.

Gallery build now over-fetches TMDB/Wikidata candidates by a configurable
factor: near-duplicate stills (the same photo at different crops or
resolutions) are discarded after embedding, so downloading exactly
images_per_actor left actors short of that many *distinct* embeddings.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-30 17:32:08 +02:00

47 lines
1.9 KiB
Python

"""Shared loader for the sae_embed nanobind module (SCRFD + ArcFace).
sae_embed.FaceEmbedder loads both ONNX sessions once and exposes an
embed(path) -> FaceResult method, avoiding the per-process model reload cost
of spawning the embed_faces CLI binary for every image.
"""
import sys
from pathlib import Path
def load_embedder(build_dir: str, models_dir: str, arcface: str | None = None,
conf: float = 0.5, nms: float = 0.4, max_side: int = 500):
"""Import sae_embed from build_dir and construct a FaceEmbedder.
Exits with a clear error if the module or models are missing — there is
no subprocess fallback.
"""
build_path = Path(build_dir).resolve()
sys.path.insert(0, str(build_path))
try:
import sae_embed
except ImportError as e:
sys.exit(
f"sae_embed module not found in {build_path}: {e}\n"
f"Build it first: cmake --build {build_dir} --target sae_embed"
)
models_path = Path(models_dir)
detector_path = str(models_path / "scrfd_500m_bnkps.onnx")
arcface_path = arcface if arcface else str(models_path / "arcface_w600k_r50.onnx")
for model, name in [(detector_path, "SCRFD"), (arcface_path, "ArcFace")]:
if not Path(model).is_file():
sys.exit(f"{name} model not found: {model}\nRun: bash scripts/download_models.sh")
# A TRT-backend build cannot load .onnx; it needs pre-built engines from
# scripts/build_trt_engines.sh. Pass them when present (ignored by ORT).
trt = Path(models_path).parent / "trt_cache"
det_engine = trt / "scrfd.scrfd_500m_bnkps.640.fp16.engine"
arc_engine = trt / f"arcface.{Path(arcface_path).stem}.b4.fp16.engine"
return sae_embed.FaceEmbedder(
detector_path, arcface_path, conf, nms, max_side,
str(det_engine) if det_engine.is_file() else "",
str(arc_engine) if arc_engine.is_file() else "",
)