first version
This commit is contained in:
141
sam2-backend/registry.py
Normal file
141
sam2-backend/registry.py
Normal file
@@ -0,0 +1,141 @@
|
||||
"""Model registry. Reads models.yaml, lazy-loads adapters, caches instances.
|
||||
|
||||
Only one adapter is resident at a time by default — switching models unloads
|
||||
the previous one to keep GPU memory bounded. Set SAM_BACKEND_KEEP_ALL=1 to
|
||||
disable eviction (useful if you have VRAM headroom and switch often).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import yaml
|
||||
|
||||
from adapters import ADAPTERS, BaseAdapter, PredictResult
|
||||
|
||||
log = logging.getLogger("sam2-backend.registry")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelSpec:
|
||||
id: str
|
||||
name: str
|
||||
family: str
|
||||
config: str
|
||||
checkpoint: str
|
||||
# If true, `checkpoint` is treated as an opaque identifier (e.g. a
|
||||
# HuggingFace model id like "facebook/sam3") rather than a file on disk.
|
||||
# Availability check is skipped for remote entries.
|
||||
remote: bool = False
|
||||
available: bool = False
|
||||
# Resolved absolute paths for logging / responses.
|
||||
checkpoint_abs: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Registry:
|
||||
specs: dict[str, ModelSpec] = field(default_factory=dict)
|
||||
default_id: Optional[str] = None
|
||||
device: str = "cpu"
|
||||
keep_all: bool = False
|
||||
_loaded: dict[str, BaseAdapter] = field(default_factory=dict)
|
||||
_lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def list_specs(self) -> list[ModelSpec]:
|
||||
return list(self.specs.values())
|
||||
|
||||
def resolve(self, model_id: Optional[str]) -> ModelSpec:
|
||||
if not model_id:
|
||||
if self.default_id is None:
|
||||
raise KeyError("no default model configured")
|
||||
model_id = self.default_id
|
||||
if model_id not in self.specs:
|
||||
raise KeyError(f"unknown model: {model_id}")
|
||||
return self.specs[model_id]
|
||||
|
||||
def get_adapter(self, model_id: str) -> BaseAdapter:
|
||||
"""Returns a loaded adapter for `model_id`, loading on first use."""
|
||||
with self._lock:
|
||||
if model_id in self._loaded:
|
||||
return self._loaded[model_id]
|
||||
spec = self.resolve(model_id)
|
||||
if not spec.available:
|
||||
raise FileNotFoundError(
|
||||
f"checkpoint missing: {spec.checkpoint_abs}"
|
||||
)
|
||||
cls = ADAPTERS.get(spec.family)
|
||||
if cls is None:
|
||||
raise ValueError(
|
||||
f"no adapter registered for family '{spec.family}'"
|
||||
)
|
||||
if not self.keep_all:
|
||||
# Evict whichever model is currently resident.
|
||||
for old_id, adapter in list(self._loaded.items()):
|
||||
log.info("evicting %s to make room for %s", old_id, model_id)
|
||||
adapter.unload()
|
||||
del self._loaded[old_id]
|
||||
adapter = cls()
|
||||
adapter.load(spec.config, spec.checkpoint_abs, self.device)
|
||||
self._loaded[model_id] = adapter
|
||||
return adapter
|
||||
|
||||
def predict_bbox(self, model_id: Optional[str], image_rgb, bbox_xyxy) -> tuple[ModelSpec, PredictResult]:
|
||||
spec = self.resolve(model_id)
|
||||
adapter = self.get_adapter(spec.id)
|
||||
# Serialize inference: most of these adapters aren't reentrant.
|
||||
with self._lock:
|
||||
result = adapter.predict_bbox(image_rgb, bbox_xyxy)
|
||||
return spec, result
|
||||
|
||||
|
||||
def load_registry(models_yaml: Path, device: str) -> Registry:
|
||||
with open(models_yaml, "r") as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
base_dir = models_yaml.parent
|
||||
keep_all = os.environ.get("SAM_BACKEND_KEEP_ALL", "0") == "1"
|
||||
reg = Registry(device=device, keep_all=keep_all)
|
||||
reg.default_id = data.get("default")
|
||||
|
||||
for entry in data.get("models", []):
|
||||
remote = bool(entry.get("remote", False))
|
||||
ckpt_str = entry["checkpoint"]
|
||||
if remote:
|
||||
# Opaque identifier — leave as-is, mark available so the UI
|
||||
# shows it. The adapter resolves/downloads on first use.
|
||||
ckpt_abs = ckpt_str
|
||||
available = True
|
||||
else:
|
||||
ckpt = Path(ckpt_str)
|
||||
ckpt_path = (ckpt if ckpt.is_absolute() else (base_dir / ckpt)).resolve()
|
||||
ckpt_abs = str(ckpt_path)
|
||||
available = ckpt_path.is_file()
|
||||
spec = ModelSpec(
|
||||
id=entry["id"],
|
||||
name=entry["name"],
|
||||
family=entry["family"],
|
||||
config=entry["config"],
|
||||
checkpoint=ckpt_str,
|
||||
remote=remote,
|
||||
available=available,
|
||||
checkpoint_abs=ckpt_abs,
|
||||
)
|
||||
reg.specs[spec.id] = spec
|
||||
log.info(
|
||||
"registered model %s (available=%s, ckpt=%s)",
|
||||
spec.id, spec.available, spec.checkpoint_abs,
|
||||
)
|
||||
|
||||
# Fall back the default to the first *available* model if the configured
|
||||
# default's checkpoint is missing.
|
||||
if reg.default_id and not reg.specs.get(reg.default_id, ModelSpec("", "", "", "", "")).available:
|
||||
for s in reg.specs.values():
|
||||
if s.available:
|
||||
reg.default_id = s.id
|
||||
break
|
||||
return reg
|
||||
Reference in New Issue
Block a user