from __future__ import annotations import importlib.util from functools import lru_cache from pathlib import Path import cv2 import numpy as np from platformdirs import user_cache_path from .models import Device SAM2_REPOSITORY = "facebook/sam2.1-hiera-tiny" class ShapeExtractionUnavailable(RuntimeError): pass def sam2_path() -> Path: return user_cache_path("editable-image") / "models" / "sam2.1-hiera-tiny" def sam2_installed() -> bool: target = sam2_path() return (target / "config.json").exists() and (target / "model.safetensors").exists() def sam2_dependencies_installed() -> bool: return all(importlib.util.find_spec(name) is not None for name in ("torch", "transformers")) def install_sam2() -> Path: if importlib.util.find_spec("huggingface_hub") is None: raise ShapeExtractionUnavailable( "Install the sam2 project extra before downloading the SAM 2 model" ) from huggingface_hub import snapshot_download target = sam2_path() target.mkdir(parents=True, exist_ok=True) snapshot_download( repo_id=SAM2_REPOSITORY, local_dir=target, allow_patterns=["*.json", "*.safetensors", "*.txt"], ) if not sam2_installed(): raise ShapeExtractionUnavailable("The downloaded SAM 2 model is incomplete") return target def refine_with_sam2( rgb: np.ndarray, residual_mask: np.ndarray, container_mask: np.ndarray, device: str, ) -> np.ndarray: if not sam2_dependencies_installed() or not sam2_installed(): raise ShapeExtractionUnavailable( "SAM 2 is unavailable. Install the sam2 extra and run " "'editable-image models install sam2'." ) generator = mask_generator(resolve_device(Device(device))) outputs = generator(rgb, points_per_batch=32) masks = outputs.get("masks", []) selected = np.zeros_like(residual_mask, dtype=np.uint8) container_area = max(cv2.countNonZero(container_mask), 1) for raw in masks: candidate = np.asarray(raw, dtype=np.uint8) if candidate.shape != residual_mask.shape: candidate = cv2.resize( candidate, (residual_mask.shape[1], residual_mask.shape[0]), interpolation=cv2.INTER_NEAREST, ) candidate = (candidate > 0).astype(np.uint8) * 255 area = cv2.countNonZero(candidate) if area == 0 or area > container_area * 0.8: continue inside = cv2.countNonZero(cv2.bitwise_and(candidate, container_mask)) / area residual_overlap = cv2.countNonZero(cv2.bitwise_and(candidate, residual_mask)) / area if inside >= 0.8 and residual_overlap >= 0.12: selected = cv2.bitwise_or(selected, candidate) selected[container_mask == 0] = 0 return selected if cv2.countNonZero(selected) else residual_mask def resolve_device(requested: Device) -> str: import torch if requested == Device.CPU: return "cpu" if requested in {Device.AUTO, Device.CUDA} and torch.cuda.is_available(): return "cuda" if requested in {Device.AUTO, Device.MPS} and torch.backends.mps.is_available(): return "mps" if requested in {Device.CUDA, Device.MPS}: raise ShapeExtractionUnavailable(f"The requested {requested.value} device is unavailable") return "cpu" @lru_cache(maxsize=3) def mask_generator(device: str): from transformers import pipeline pipeline_device: int | str = 0 if device == "cuda" else device return pipeline( "mask-generation", model=str(sam2_path()), device=pipeline_device, )