Spaces:
Build error
Build error
Download src/editable_image/sam2.py from tonigi/make_editable_image: direct link, hf CLI and curl.
- Browser
- Download file 3.66 kB
-
https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/sam2.py
- Command line
-
hf download hf://spaces/tonigi/make_editable_image/src/editable_image/sam2.py
-
curl -L -o sam2.py https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/sam2.py
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" | |
| 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, | |
| ) | |