Files
SnapOtter/packages/ai/python/remove_bg.py
T
SnapOtter 80957f6e10 feat: improve remove background with edge smoothing, color decontamination, output formats
- Expose birefnet-hr-matting in UI (People/Ultra) and fix model defaults
  (People/Max now uses birefnet-matting for true alpha matting)
- Add output format selector (PNG/WebP/AVIF) with lossless alpha support
- Add edge smoothing post-processing (Off/Light/Medium/Strong) via
  morphological mask refinement to reduce gray halo artifacts
- Add color decontamination to remove background color spill from
  semi-transparent edge pixels
- Thread new settings through full stack: frontend -> API schema ->
  Python sidecar -> Sharp effects pipeline
- Add i18n keys for all 21 locales
- Add unit tests for new option serialization (3 tests)
- Add integration tests for new settings validation (4 tests)
2026-06-05 23:05:25 +08:00

277 lines
8.6 KiB
Python

"""Background removal using rembg with state-of-the-art BiRefNet models."""
import sys
import json
import os
def emit_progress(percent, stage):
print(json.dumps({"progress": percent, "stage": stage}), file=sys.stderr, flush=True)
def _refine_edges(image_bytes, level):
"""Morphological mask refinement to reduce gray halos on edges.
level: 1=light, 2=medium, 3=strong
"""
import cv2
import numpy as np
from PIL import Image
import io
img = Image.open(io.BytesIO(image_bytes)).convert("RGBA")
arr = np.array(img)
alpha = arr[:, :, 3]
kernel_size = 1 + level
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))
alpha = cv2.morphologyEx(alpha, cv2.MORPH_CLOSE, kernel)
sigma = 0.3 + level * 0.3
alpha = cv2.GaussianBlur(alpha, (0, 0), sigma)
arr[:, :, 3] = alpha
out = Image.fromarray(arr, "RGBA")
buf = io.BytesIO()
out.save(buf, format="PNG")
return buf.getvalue()
def _decontaminate_edges(image_bytes):
"""Remove background color spill from semi-transparent edge pixels."""
import numpy as np
from PIL import Image
import io
img = Image.open(io.BytesIO(image_bytes)).convert("RGBA")
arr = np.array(img, dtype=np.float32)
alpha = arr[:, :, 3] / 255.0
rgb = arr[:, :, :3]
bg_mask = alpha < 0.04
if not np.any(bg_mask):
return image_bytes
bg_color = np.zeros(3, dtype=np.float32)
for c in range(3):
channel = rgb[:, :, c]
bg_pixels = channel[bg_mask]
if len(bg_pixels) > 0:
bg_color[c] = np.median(bg_pixels)
edge_mask = (alpha > 0.04) & (alpha < 0.96)
if not np.any(edge_mask):
return image_bytes
a = alpha[edge_mask, np.newaxis]
fg = rgb[edge_mask]
corrected = (fg - bg_color[np.newaxis, :] * (1.0 - a)) / np.maximum(a, 0.01)
corrected = np.clip(corrected, 0, 255)
rgb[edge_mask] = corrected
arr[:, :, :3] = rgb
result = np.clip(arr, 0, 255).astype(np.uint8)
out = Image.fromarray(result, "RGBA")
buf = io.BytesIO()
out.save(buf, format="PNG")
return buf.getvalue()
ALLOWED_MODELS = {
"u2net",
"isnet-general-use",
"bria-rmbg",
"birefnet-general-lite",
"birefnet-portrait",
"birefnet-general",
"birefnet-matting",
"birefnet-hr-matting",
}
_matting_registered = False
def _register_matting_session(sessions_class):
"""Register the BiRefNet-matting ONNX session for Ultra quality mode."""
global _matting_registered
if _matting_registered:
return
_matting_registered = True
import os
import pooch
from rembg.sessions.birefnet_general import BiRefNetSessionGeneral
class BiRefNetMattingSession(BiRefNetSessionGeneral):
@classmethod
def download_models(cls, *args, **kwargs):
fname = f"{cls.name(*args, **kwargs)}.onnx"
pooch.retrieve(
"https://github.com/ZhengPeng7/BiRefNet/releases/download/v1/BiRefNet-matting-epoch_100.onnx",
None, # Skip checksum for GitHub release assets
fname=fname,
path=cls.u2net_home(*args, **kwargs),
progressbar=True,
)
return os.path.join(cls.u2net_home(*args, **kwargs), fname)
@classmethod
def name(cls, *args, **kwargs):
return "birefnet-matting"
sessions_class.append(BiRefNetMattingSession)
_hr_matting_registered = False
def _register_hr_matting_session(sessions_class):
"""Register the BiRefNet HR-matting ONNX session for 2048x2048 high-res matting."""
global _hr_matting_registered
if _hr_matting_registered:
return
_hr_matting_registered = True
import os
import numpy as np
import pooch
from PIL import Image
from rembg.sessions.birefnet_general import BiRefNetSessionGeneral
class BiRefNetHRMattingSession(BiRefNetSessionGeneral):
@classmethod
def download_models(cls, *args, **kwargs):
fname = f"{cls.name(*args, **kwargs)}.onnx"
pooch.retrieve(
"https://github.com/ZhengPeng7/BiRefNet/releases/download/v1/BiRefNet_HR-matting-epoch_135.onnx",
None,
fname=fname,
path=cls.u2net_home(*args, **kwargs),
progressbar=True,
)
return os.path.join(cls.u2net_home(*args, **kwargs), fname)
@classmethod
def name(cls, *args, **kwargs):
return "birefnet-hr-matting"
def predict(self, img, *args, **kwargs):
ort_outs = self.inner_session.run(
None,
self.normalize(
img, (0.485, 0.456, 0.406), (0.229, 0.224, 0.225), (2048, 2048)
),
)
pred = ort_outs[0][:, 0, :, :]
ma = np.max(pred)
mi = np.min(pred)
denom = ma - mi
pred = (pred - mi) / denom if denom > 0 else pred * 0
pred = np.squeeze(pred)
mask = Image.fromarray((pred * 255).astype("uint8"), mode="L")
mask = mask.resize(img.size, Image.LANCZOS)
return [mask]
sessions_class.append(BiRefNetHRMattingSession)
def main():
input_path = sys.argv[1]
output_path = sys.argv[2]
settings = json.loads(sys.argv[3]) if len(sys.argv) > 3 else {}
model = settings.get("model", "birefnet-general-lite")
if model not in ALLOWED_MODELS:
model = "birefnet-general-lite"
# Redirect stdout to stderr so library download/progress output
# cannot contaminate our JSON result on stdout.
stdout_fd = os.dup(1)
os.dup2(2, 1)
try:
from rembg import remove, new_session
from rembg.sessions import sessions_class
from gpu import onnx_providers
# Register BiRefNet-matting (Ultra quality) if not already present
_register_matting_session(sessions_class)
_register_hr_matting_session(sessions_class)
emit_progress(10, "Loading model")
providers, device = onnx_providers()
try:
session = new_session(model, providers=providers)
except Exception as e:
if "CUDAExecutionProvider" in providers:
from gpu import emit_info
emit_info(f"GPU session failed ({e}), falling back to CPU")
session = new_session(model, providers=["CPUExecutionProvider"])
device = "cpu"
else:
raise
emit_progress(25, "Model loaded")
with open(input_path, "rb") as f:
input_data = f.read()
emit_progress(30, "Analyzing image")
use_alpha_matting = device != "cpu"
try:
output_data = remove(
input_data,
session=session,
alpha_matting=use_alpha_matting,
alpha_matting_foreground_threshold=240,
alpha_matting_background_threshold=10,
)
except Exception as e:
if use_alpha_matting:
emit_progress(35, "Retrying without alpha matting")
output_data = remove(input_data, session=session, alpha_matting=False)
else:
raise RuntimeError(
f"Background removal failed: {e}"
) from e
emit_progress(80, "Background removed")
edge_refine = settings.get("edgeRefine", 0)
decontaminate = settings.get("decontaminate", False)
if edge_refine and edge_refine > 0:
emit_progress(85, "Refining edges")
output_data = _refine_edges(output_data, int(edge_refine))
if decontaminate:
emit_progress(90, "Removing color spill")
output_data = _decontaminate_edges(output_data)
# Always return transparent PNG. All background compositing
# (solid color, gradient, blur, shadow) is handled by Node.js/Sharp.
emit_progress(95, "Saving result")
with open(output_path, "wb") as f:
f.write(output_data)
result = json.dumps({"success": True, "model": model, "device": device})
except ImportError as e:
print(f"[remove-bg] Import failed: {e}", file=sys.stderr, flush=True)
result = json.dumps(
{
"success": False,
"error": f"rembg import failed: {e}",
}
)
except Exception as e:
result = json.dumps({"success": False, "error": str(e)})
# Restore original stdout and write only our JSON result
os.dup2(stdout_fd, 1)
os.close(stdout_fd)
sys.stdout.write(result + "\n")
sys.stdout.flush()
if __name__ == "__main__":
main()