"""SAM3 text-prompted detection, wrapped for reuse across labeling jobs. The model is expensive to build (weights come from the gated HuggingFace repo `facebook/sam3`), so it is loaded once per process and kept resident. The important performance detail: `Sam3Processor.set_image()` runs the vision backbone, while `set_text_prompt()` only runs the (much cheaper) grounding head against the cached `backbone_out`. So for an N-prompt job we call `set_image` once per image and loop the prompts over that same state. """ import os import sys import threading from dataclasses import dataclass, field from typing import List, Optional import numpy as np import torch from PIL import Image # Fix python import path masking issue where sam3 is imported as a namespace package sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "sam3"))) from sam3.model.sam3_image_processor import Sam3Processor from sam3.model_builder import build_sam3_image_model @dataclass class Detection: """One labeled instance in one image.""" class_id: int class_name: str = "" score: float = 0.0 box: List[float] = field(default_factory=list) # xyxy in pixels mask: Optional[np.ndarray] = None # bool array, (H, W) at original image size class Sam3Engine: def __init__(self, checkpoint_path: Optional[str] = None): # SAM3 is CUDA-only in practice: `PositionEmbeddingSine` precomputes its # tables with a hardcoded `device="cuda"`, so a CPU run dies deep inside # the backbone with an unrelated-looking error. Fail here instead, where # the message can say something useful. if not torch.cuda.is_available(): raise RuntimeError( "SAM3 requires a CUDA GPU. No GPU is visible to torch — check " "`nvidia-smi`, that CUDA_VISIBLE_DEVICES isn't set to empty, and " "that this venv has a CUDA build of torch installed." ) self.device = "cuda" self.autocast_dtype = torch.float16 # `enable_inst_interactivity=True` builds a SAM1-style click predictor, # but in this vendored copy its `image_encoder` is None and its expected # feature sizes (288/144/72) don't match the 1008px image pipeline, so # `predictor.set_image()` always fails. It costs ~0.4 GB for nothing, so # it stays off. Box exemplars cover the interactive use case instead. self.supports_tap = False self.model = build_sam3_image_model( device=self.device, checkpoint_path=checkpoint_path, load_from_HF=checkpoint_path is None, ) self.processor = Sam3Processor(self.model, device=self.device) def detect(self, image: Image.Image, prompts: List[str], threshold: float, thresholds: Optional[List[float]] = None) -> List[Detection]: """Run every prompt against one image; prompt index becomes the class id. `thresholds` overrides the confidence per prompt (REQ-181); still one `set_image` for the whole call — only the grounding head sees the change.""" processor = Sam3Processor(self.model, device=self.device) processor.confidence_threshold = threshold detections: List[Detection] = [] with torch.autocast(self.device, dtype=self.autocast_dtype): state = processor.set_image(image) for class_id, prompt in enumerate(prompts): if thresholds is not None and class_id < len(thresholds): processor.confidence_threshold = thresholds[class_id] output = processor.set_text_prompt(prompt=prompt, state=state) masks, boxes, scores = output["masks"], output["boxes"], output["scores"] if masks.shape[0] == 0: continue # Pull off the GPU immediately: the next prompt overwrites these # tensors, and full-resolution masks are the memory hog here. masks_np = masks.squeeze(1).to(torch.uint8).cpu().numpy().astype(bool) boxes_np = boxes.float().cpu().numpy() scores_np = scores.float().cpu().numpy() for i in range(masks_np.shape[0]): detections.append( Detection( class_id=class_id, class_name=prompt, score=float(scores_np[i]), box=[float(v) for v in boxes_np[i]], mask=masks_np[i], ) ) del state return detections def detect_with_exemplars( self, image: Image.Image, prompts: List[str], threshold: float, exemplar_index: int, exemplars: List[dict], thresholds: Optional[List[float]] = None, ) -> List[Detection]: """`detect()`, but one prompt also carries drawn box exemplars (REQ-172). Still one `set_image` for the whole call. The prompt set is reset before every class because `state["geometric_prompt"]` survives `set_text_prompt` — without the reset, one class's boxes would leak into the next class. """ processor = Sam3Processor(self.model, device=self.device) processor.confidence_threshold = threshold detections: List[Detection] = [] with torch.autocast(self.device, dtype=self.autocast_dtype): state = processor.set_image(image) for class_id, prompt in enumerate(prompts): if thresholds is not None and class_id < len(thresholds): processor.confidence_threshold = thresholds[class_id] processor.reset_all_prompts(state) output = processor.set_text_prompt(prompt=prompt, state=state) if class_id == exemplar_index: for exemplar in exemplars: output = processor.add_geometric_prompt( box=exemplar["box"], label=bool(exemplar.get("positive", True)), state=state, ) detections.extend(self._collect(output, class_id, prompt)) del state return detections # ---- interactive / exemplar prompting ------------------------------ def open_state(self, image: Image.Image): """Run the vision backbone once and hand back the reusable state.""" with torch.autocast(self.device, dtype=self.autocast_dtype): return self.processor.set_image(image) def apply_prompts( self, state, threshold: float, text: Optional[str] = None, exemplars: Optional[List[dict]] = None, ) -> List[Detection]: """Re-run grounding for this image from scratch with the given prompts. Exemplars are boxes in normalized cxcywh with a positive/negative flag. The prompt set is always replayed from empty because SAM3 only supports appending geometric prompts — that's how undo is implemented. """ self.processor.confidence_threshold = threshold exemplars = exemplars or [] with torch.autocast(self.device, dtype=self.autocast_dtype): self.processor.reset_all_prompts(state) output = None if text: output = self.processor.set_text_prompt(prompt=text, state=state) for exemplar in exemplars: output = self.processor.add_geometric_prompt( box=exemplar["box"], label=bool(exemplar.get("positive", True)), state=state, ) if output is None: return [] return self._collect(output, class_id=0, class_name=text or "visual") def segment_at( self, image: Image.Image, points: Optional[List[List[float]]] = None, labels: Optional[List[int]] = None, box: Optional[List[float]] = None, ) -> Optional[Detection]: """Tap-to-segment: one point (or box) in pixels -> that object's mask.""" if not self.supports_tap: return None predictor = self.model.inst_interactive_predictor predictor.set_image(np.array(image)) masks, scores, _ = predictor.predict( point_coords=np.array(points, dtype=np.float32) if points else None, point_labels=np.array(labels, dtype=np.int32) if labels else None, box=np.array(box, dtype=np.float32) if box else None, multimask_output=True, ) if masks.shape[0] == 0: return None best = int(np.argmax(scores)) mask = masks[best].astype(bool) ys, xs = np.where(mask) if xs.size == 0: return None return Detection( class_id=0, class_name="tap", score=float(scores[best]), box=[float(xs.min()), float(ys.min()), float(xs.max()), float(ys.max())], mask=mask, ) def _collect(self, output, class_id: int, class_name: str) -> List[Detection]: masks, boxes, scores = output["masks"], output["boxes"], output["scores"] if masks.shape[0] == 0: return [] masks_np = masks.squeeze(1).to(torch.uint8).cpu().numpy().astype(bool) boxes_np = boxes.float().cpu().numpy() scores_np = scores.float().cpu().numpy() return [ Detection( class_id=class_id, class_name=class_name, score=float(scores_np[i]), box=[float(v) for v in boxes_np[i]], mask=masks_np[i], ) for i in range(masks_np.shape[0]) ] _engine: Optional[Sam3Engine] = None _engine_lock = threading.Lock() def get_engine() -> Sam3Engine: """Build the model on first use, then hand out the same instance.""" global _engine from backend import hardware with _engine_lock: if _engine is None: needed = hardware.SAM3_RESIDENT_GB + hardware.SAM3_HEADROOM_GB free = hardware.free_vram_gb() if free < needed: raise RuntimeError( f"SAM3 needs ~{needed:.1f} GB free but only {free:.1f} GB is available. " "Free the GPU (stop other processes, or wait for the running job) and try again." ) try: _engine = Sam3Engine() except (ImportError, RuntimeError) as exc: curr_free = hardware.free_vram_gb() raise RuntimeError( f"{exc} (Available VRAM: {curr_free:.1f} GB)" ) from exc return _engine def engine_is_loaded() -> bool: return _engine is not None def release_engine() -> bool: """Drop the model and free its VRAM (REQ-065). SAM3 holds ~3.4 GB resident. On a 6 GB card that is most of the memory a training run needs, so the two must never be loaded at once. The next job that needs SAM3 rebuilds it from the local cache in about 12 seconds. """ global _engine import gc with _engine_lock: if _engine is None: return False _engine = None gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() return True