"""SAM3 adapter. Mirrors the loading pattern from `files/sam3_server.py` in this repo. Two modes are supported, selected via the model spec's `config` field in models.yaml: config: image_predictor — uses `sam3.model_builder.build_sam3_image_model` with a local `.pt` checkpoint. Treats the user's bbox as a visual exemplar and returns the mask whose predicted box has highest IoU with the user box (single-object annotation workflow). config: tracker — uses HuggingFace `transformers.Sam3TrackerModel` (SAM2-style PVS on images). `checkpoint` here is a HF model ID like "facebook/sam3" rather than a local .pt. Requires `transformers` installed and network access for first download. Both paths return a single polygon + confidence for the best mask, matching the adapter contract the server's `/predict` route expects. """ 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.sam3") class Sam3Adapter(BaseAdapter): def __init__(self) -> None: self._mode: str = "image_predictor" self._device: str = "cpu" # image_predictor path self._image_model: Optional[object] = None self._image_processor: Optional[object] = None # tracker path self._tracker_model: Optional[object] = None self._tracker_processor: Optional[object] = None def load(self, config: str, checkpoint: str, device: str) -> None: self._device = device mode = (config or "image_predictor").strip().lower() self._mode = mode if mode == "tracker": self._load_tracker(checkpoint, device) elif mode in ("image_predictor", "image", "exemplar"): self._load_image_predictor(checkpoint, device) else: raise ValueError( f"unknown SAM3 config mode '{mode}'. " f"expected 'image_predictor' or 'tracker'." ) def _load_image_predictor(self, checkpoint: str, device: str) -> None: # Deferred imports so the server can start even without sam3 installed # (the model stays unavailable until first use). from sam3.model_builder import build_sam3_image_model from sam3.model.sam3_image_processor import Sam3Processor log.info("loading SAM3 image predictor: ckpt=%s device=%s", checkpoint, device) model = ( build_sam3_image_model(checkpoint_path=checkpoint) if checkpoint else build_sam3_image_model() ) self._image_model = model.to(device).eval() self._image_processor = Sam3Processor(self._image_model) log.info("SAM3 image predictor loaded") def _load_tracker(self, hf_id: str, device: str) -> None: from transformers import Sam3TrackerModel, Sam3TrackerProcessor # type: ignore model_id = hf_id or "facebook/sam3" log.info("loading SAM3 tracker: hf_id=%s device=%s", model_id, device) self._tracker_processor = Sam3TrackerProcessor.from_pretrained(model_id) self._tracker_model = ( Sam3TrackerModel.from_pretrained(model_id).to(device).eval() ) log.info("SAM3 tracker loaded") def unload(self) -> None: self._image_model = None self._image_processor = None self._tracker_model = None self._tracker_processor = None try: import torch if torch.cuda.is_available(): torch.cuda.empty_cache() except Exception: pass log.info("SAM3 unloaded") def predict_bbox( self, image_rgb: np.ndarray, bbox_xyxy: tuple[int, int, int, int] ) -> PredictResult: import torch from PIL import Image pil_image = Image.fromarray(image_rgb) # SAM3 weights are loaded in BFloat16 by `build_sam3_image_model`, but # the processor hands the model Float32 image tensors → matmul dtype # mismatch. Autocast does the input cast for us. Tracker path goes # through HF transformers which already handles dtype internally, but # autocasting is harmless there too. device_type = "cuda" if self._device.startswith("cuda") else "cpu" with torch.inference_mode(): with torch.autocast(device_type=device_type, dtype=torch.bfloat16): if self._mode == "tracker": return self._predict_tracker(pil_image, bbox_xyxy) return self._predict_image_predictor(pil_image, bbox_xyxy) # -------------------------------------------------------------------- # Tracker path — HF transformers `Sam3Tracker` # -------------------------------------------------------------------- def _predict_tracker(self, image, bbox_xyxy) -> PredictResult: if self._tracker_model is None or self._tracker_processor is None: raise RuntimeError("SAM3 tracker not loaded") x1, y1, x2, y2 = [float(v) for v in bbox_xyxy] inputs = self._tracker_processor( images=image, input_boxes=[[[x1, y1, x2, y2]]], # batch[1] × obj[1] × 4 return_tensors="pt", ).to(self._device) outputs = self._tracker_model(**inputs, multimask_output=True) # `post_process_masks` lost `reshaped_input_sizes` in newer # transformers — only `masks` and `original_sizes` are accepted now. # Cast to float32 before .numpy() — autocast leaves outputs in # bfloat16, which NumPy can't ingest. masks = self._tracker_processor.post_process_masks( outputs.pred_masks.float().cpu(), original_sizes=inputs["original_sizes"].cpu(), )[0] # (num_masks, H, W) scores = outputs.iou_scores[0].float().cpu().numpy().ravel() masks_np = masks.float().numpy().astype(bool) if masks.dtype != bool else masks.numpy() n_masks = masks_np.shape[0] if n_masks == 0: return PredictResult(polygon=[], score=0.0) # Mask and score counts can disagree (`multimask_output=True` returns 3 # IoU scores even when the head only emits 1 mask, e.g. when a # `sam3_video` checkpoint is loaded into `Sam3TrackerModel`). Pick the # best within whichever count is smaller. valid = scores[:n_masks] if scores.size >= n_masks else scores best = int(np.argmax(valid)) if valid.size > 0 else 0 best_score = float(valid[best]) if valid.size > 0 else 0.0 return PredictResult( polygon=_mask_to_polygon(masks_np[best]), score=best_score, ) # -------------------------------------------------------------------- # Image-predictor path — box-as-exemplar, pick highest-IoU result # -------------------------------------------------------------------- def _predict_image_predictor(self, image, bbox_xyxy) -> PredictResult: if self._image_processor is None: raise RuntimeError("SAM3 image predictor not loaded") x1, y1, x2, y2 = [float(v) for v in bbox_xyxy] state = self._image_processor.set_image(image) # type: ignore[attr-defined] out = self._image_processor.set_box_prompt( # type: ignore[attr-defined] state=state, boxes=[[x1, y1, x2 - x1, y2 - y1]], # xywh per SAM3 API box_labels=[1], ) masks = out.get("masks", []) boxes = out.get("boxes", []) scores = out.get("scores", []) if len(masks) == 0: return PredictResult(polygon=[], score=0.0) # Pick the proposal whose predicted bbox overlaps the user bbox most. user = (x1, y1, x2, y2) ious = [_iou(user, tuple(float(v) for v in b)) for b in boxes] idx = int(np.argmax(ious)) mask_np = np.asarray(masks[idx]).astype(bool) return PredictResult( polygon=_mask_to_polygon(mask_np), score=float(scores[idx]), ) # -------------------------------------------------------------------------- # Helpers # -------------------------------------------------------------------------- def _iou( a: tuple[float, float, float, float], b: tuple[float, float, float, float], ) -> float: ax1, ay1, ax2, ay2 = a bx1, by1, bx2, by2 = b ix1, iy1 = max(ax1, bx1), max(ay1, by1) ix2, iy2 = min(ax2, bx2), min(ay2, by2) iw = max(0.0, ix2 - ix1) ih = max(0.0, iy2 - iy1) inter = iw * ih ua = (ax2 - ax1) * (ay2 - ay1) + (bx2 - bx1) * (by2 - by1) - inter return inter / ua if ua > 0 else 0.0