Files
reTraining/backend/sam3_engine.py
T

282 lines
11 KiB
Python

"""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) -> List[Detection]:
"""Run every prompt against one image; prompt index becomes the class id."""
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):
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],
) -> 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):
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