- REQ-181: class_params {name: {threshold?, iou_threshold?, min_box_frac?}}
on /preview and /autolabel, per-class override table in both modals;
empty overrides take the unchanged global path
- REQ-180: x button on each review sidebar class row clears that class on
the current frame only via bulk-delete, no confirmation
- includes REQ-178 empty date-folder cycle fix (archive_index.py)
291 lines
11 KiB
Python
291 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,
|
|
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
|