Spaces:
Running on Zero
Running on Zero
| """WP-22: click-to-segment — the user taps an object, SAM supplies the precision. | |
| The brush asked users to be precise with their hand (shading four paintings with | |
| a mouse; hopeless on a phone), then guided-filter propagation guessed the edges. | |
| Click-to-segment inverts the contract: the user supplies INTENT (one tap: "this | |
| thing, it's a painting, Scene 2") and a promptable segmenter supplies the pixels. | |
| Verified live on photo 128 (2026-07-19): SlimSAM-uniform-50 (28M params) returns | |
| whole-print masks with real edges THROUGH the other exposure's ghosts, IoU | |
| 0.79–0.93, 0.7 s one-time image embed + ~33 ms per click on CPU. | |
| Everything is lazy and fail-soft: no transformers / no weights → point_mask | |
| returns None and the UI falls back to the brush. The photo never leaves the | |
| process — segmentation is local, free, and private. | |
| """ | |
| from __future__ import annotations | |
| from typing import List, Optional, Tuple | |
| import numpy as np | |
| # Small enough for Space cold-starts (~28 MB weights), good enough for tap-to-mask | |
| # on clear objects; the brush remains for everything else. | |
| SEGMENT_MODEL: str = "Zigeng/SlimSAM-uniform-50" | |
| _MODEL = None # (model, processor) after first successful load | |
| _LOAD_FAILED: Optional[str] = None | |
| # Single-slot embedding cache: (key, embeddings, original PIL size). One photo is | |
| # worked on at a time; re-embedding on photo switch costs ~1 s. | |
| _EMBED: dict = {} | |
| def _load(): | |
| global _MODEL, _LOAD_FAILED | |
| if _MODEL is not None or _LOAD_FAILED is not None: | |
| return _MODEL | |
| try: | |
| from transformers import SamModel, SamProcessor | |
| model = SamModel.from_pretrained(SEGMENT_MODEL) | |
| model.eval() | |
| _MODEL = (model, SamProcessor.from_pretrained(SEGMENT_MODEL)) | |
| except Exception as exc: # missing dep, no network for weights, OOM… | |
| _LOAD_FAILED = str(exc) | |
| return _MODEL | |
| def _image_key(rgb: np.ndarray) -> tuple: | |
| """Cheap content key: shape + a sparse pixel checksum (no full-image hash).""" | |
| flat = rgb.reshape(-1) | |
| stride = max(1, flat.size // 4096) | |
| return (rgb.shape, float(flat[::stride].sum())) | |
| def point_mask( | |
| rgb: np.ndarray, points: List[Tuple[float, float]] | |
| ) -> Optional[Tuple[np.ndarray, float]]: | |
| """Mask for the object under the given positive click points. | |
| ``rgb`` float [0,1] HxWx3 in the SAME geometry the user clicked (the working | |
| frame); ``points`` are (x, y) pixel coords, several clicks refining one | |
| object. Returns (bool HxW mask, iou confidence), or None when the model is | |
| unavailable (caller falls back to the brush) or the points are empty. | |
| """ | |
| if not points: | |
| return None | |
| loaded = _load() | |
| if loaded is None: | |
| return None | |
| model, processor = loaded | |
| import torch | |
| from PIL import Image | |
| pil = Image.fromarray((np.clip(rgb, 0, 1) * 255).astype(np.uint8)) | |
| key = _image_key(rgb) | |
| if _EMBED.get("key") != key: | |
| inputs = processor(pil, return_tensors="pt") | |
| with torch.no_grad(): | |
| emb = model.get_image_embeddings(inputs["pixel_values"]) | |
| _EMBED.update(key=key, emb=emb) | |
| inp = processor(pil, input_points=[[list(p) for p in points]], return_tensors="pt") | |
| inp.pop("pixel_values") | |
| with torch.no_grad(): | |
| out = model(**inp, image_embeddings=_EMBED["emb"], multimask_output=True) | |
| masks = processor.image_processor.post_process_masks( | |
| out.pred_masks.cpu(), inp["original_sizes"].cpu(), inp["reshaped_input_sizes"].cpu() | |
| )[0][0] | |
| scores = out.iou_scores[0, 0].cpu().numpy() | |
| best = int(scores.argmax()) | |
| return masks[best].numpy().astype(bool), float(scores[best]) | |
| def load_error() -> Optional[str]: | |
| """Why the segmenter is unavailable (None while unloaded or fine).""" | |
| return _LOAD_FAILED | |