Files
Vase-Matcher/pipeline/bg_removal.py
2026-08-04 17:09:29 +05:30

148 lines
5.0 KiB
Python

"""
Background removal via rembg (BiRefNet lite), disk-cached by content hash
so a re-upload of the same image never re-runs the model.
"""
import gc
import hashlib
import io
import logging
import os
import cv2
import numpy as np
from PIL import Image
from rembg import remove, new_session
import config
logger = logging.getLogger(__name__)
_session = None
def get_session():
"""Lazily create the rembg session once and reuse it for every request."""
global _session
if _session is None:
try:
_session = new_session(config.REMBG_MODEL_NAME,
providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
except Exception:
logger.exception(
"Failed to create rembg session with CUDAExecutionProvider, "
"falling back to CPU-only"
)
_session = new_session(config.REMBG_MODEL_NAME, providers=["CPUExecutionProvider"])
try:
providers = _session.inner_session.get_providers()
logger.info("rembg session ready, using providers: %s", providers)
if "CUDAExecutionProvider" not in providers:
logger.warning(
"rembg is running on CPU (no CUDAExecutionProvider) -- "
"background removal will be much slower. Check that "
"onnxruntime-gpu is installed and the CUDA driver is visible."
)
except AttributeError:
pass
return _session
def unload_session():
"""Drops the rembg/onnxruntime session so its (possibly CUDA-backed)
memory is freed. It's a lazy singleton (see get_session above), so the
next call to remove_background_bytes/_file simply recreates it. Used to
make room for SAM-based flower counting on an 8GB card."""
global _session
freed = _session is not None
_session = None
if freed:
gc.collect()
return freed
def resize_max_dim(pil_img: Image.Image, max_dim: int) -> Image.Image:
"""Downscale in place-equivalent fashion so neither side exceeds max_dim.
No-op (returns the same image) if already within bounds -- this is the
single choke point that keeps memory/time bounded for every model in the
pipeline (bg removal, SIFT/ORB, SuperPoint, LoFTR), regardless of how
large the original upload was."""
w, h = pil_img.size
scale = max_dim / max(w, h)
if scale >= 1.0:
return pil_img
new_size = (max(1, int(w * scale)), max(1, int(h * scale)))
return pil_img.resize(new_size, Image.LANCZOS)
def hash_bytes(data: bytes) -> str:
return hashlib.sha1(data).hexdigest()[:20]
def remove_background_bytes(image_bytes: bytes, cache_dir: str):
"""
Runs background removal on raw image bytes (an upload), caching the
result by content hash. Returns (rgba_bgra_ndarray, cache_key).
"""
os.makedirs(cache_dir, exist_ok=True)
key = hash_bytes(image_bytes)
cache_path = os.path.join(cache_dir, key + ".png")
if os.path.exists(cache_path):
img = cv2.imread(cache_path, cv2.IMREAD_UNCHANGED)
if img is not None and img.ndim == 3 and img.shape[2] == 4:
return img, key
pil_img = Image.open(io.BytesIO(image_bytes)).convert("RGB")
orig_size = pil_img.size
pil_img = resize_max_dim(pil_img, config.MAX_IMAGE_DIM)
if pil_img.size != orig_size:
logger.info("Resized upload %s -> %s before processing", orig_size, pil_img.size)
result = remove(pil_img, session=get_session())
result.save(cache_path)
rgba = cv2.cvtColor(np.array(result), cv2.COLOR_RGBA2BGRA)
return rgba, key
def remove_background_file(src_path: str, cache_dir: str):
"""Same as remove_background_bytes, but caches by original filename --
used for the fixed template set, which doesn't change between requests."""
os.makedirs(cache_dir, exist_ok=True)
base_name = os.path.splitext(os.path.basename(src_path))[0]
cache_path = os.path.join(cache_dir, base_name + ".png")
if os.path.exists(cache_path):
img = cv2.imread(cache_path, cv2.IMREAD_UNCHANGED)
if img is not None and img.ndim == 3 and img.shape[2] == 4:
return img
pil_img = Image.open(src_path).convert("RGB")
pil_img = resize_max_dim(pil_img, config.MAX_IMAGE_DIM)
result = remove(pil_img, session=get_session())
result.save(cache_path)
return cv2.cvtColor(np.array(result), cv2.COLOR_RGBA2BGRA)
def create_mask(alpha):
_, mask = cv2.threshold(alpha, config.ALPHA_THRESHOLD, 255, cv2.THRESH_BINARY)
kernel = np.ones((3, 3), np.uint8)
mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel)
mask = cv2.erode(mask, kernel, iterations=config.ERODE_ITER)
return mask
def split_rgba(rgba):
"""rgba: BGRA ndarray -> (bgr, alpha, mask)"""
if rgba.shape[2] == 4:
bgr = rgba[:, :, :3]
alpha = rgba[:, :, 3]
else:
bgr = rgba
alpha = np.ones(bgr.shape[:2], dtype=np.uint8) * 255
mask = create_mask(alpha)
return bgr, alpha, mask