first version
This commit is contained in:
82
sam2-backend/adapters/sam2_adapter.py
Normal file
82
sam2-backend/adapters/sam2_adapter.py
Normal file
@@ -0,0 +1,82 @@
|
||||
"""SAM2 adapter — wraps `sam2.sam2_image_predictor.SAM2ImagePredictor`.
|
||||
|
||||
The SAM2 repo uses Hydra to resolve config files relative to the `sam2`
|
||||
package. We rely on that resolution: pass `config="configs/sam2.1/..."` and
|
||||
it works as long as the `sam2` package is importable (pip-installed or on
|
||||
PYTHONPATH).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .base import BaseAdapter, PredictResult
|
||||
from .utils import mask_to_polygon as _mask_to_polygon
|
||||
|
||||
log = logging.getLogger("sam2-backend.sam2")
|
||||
|
||||
|
||||
class Sam2Adapter(BaseAdapter):
|
||||
def __init__(self) -> None:
|
||||
self._predictor: Optional[object] = None
|
||||
self._device: str = "cpu"
|
||||
|
||||
def load(self, config: str, checkpoint: str, device: str) -> None:
|
||||
# Imports are deferred so the server can start even without the SAM2
|
||||
# deps installed (the model stays in `available: false` until loaded).
|
||||
from sam2.build_sam import build_sam2
|
||||
from sam2.sam2_image_predictor import SAM2ImagePredictor
|
||||
|
||||
log.info("loading SAM2: config=%s ckpt=%s device=%s",
|
||||
config, checkpoint, device)
|
||||
model = build_sam2(config, checkpoint, device=device)
|
||||
self._predictor = SAM2ImagePredictor(model)
|
||||
self._device = device
|
||||
log.info("SAM2 loaded")
|
||||
|
||||
def unload(self) -> None:
|
||||
if self._predictor is None:
|
||||
return
|
||||
self._predictor = None
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
log.info("SAM2 unloaded")
|
||||
|
||||
def predict_bbox(
|
||||
self, image_rgb: np.ndarray, bbox_xyxy: tuple[int, int, int, int]
|
||||
) -> PredictResult:
|
||||
if self._predictor is None:
|
||||
raise RuntimeError("SAM2 adapter not loaded")
|
||||
|
||||
import torch
|
||||
|
||||
# set_image re-embeds on every call; that's fine for an interactive
|
||||
# one-box-per-request workflow. Batch prompts would want caching.
|
||||
inference_ctx = (
|
||||
torch.inference_mode()
|
||||
if self._device.startswith("cuda") or self._device == "cpu"
|
||||
else torch.no_grad()
|
||||
)
|
||||
with inference_ctx:
|
||||
self._predictor.set_image(image_rgb) # type: ignore[attr-defined]
|
||||
box = np.array(list(bbox_xyxy), dtype=np.float32)
|
||||
masks, scores, _ = self._predictor.predict( # type: ignore[attr-defined]
|
||||
box=box,
|
||||
multimask_output=True,
|
||||
)
|
||||
|
||||
# multimask_output returns 3 candidates — take the highest-scoring one.
|
||||
best_idx = int(np.argmax(scores))
|
||||
best_mask = masks[best_idx]
|
||||
best_score = float(scores[best_idx])
|
||||
polygon = _mask_to_polygon(best_mask)
|
||||
return PredictResult(polygon=polygon, score=best_score)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user