tonigi's picture
Add complete Gradio interface
3c2cd23 verified
Raw History Blame Contribute Delete
3.66 kB
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,
)