Add BEN2 and BiRefNet-HR as selectable Remove Background models
BEN2 becomes the new default local backend (clean cutouts, strong on hair/fur edges), with BiRefNet-HR available as a high-res/print alternate and U2Net kept as the lightweight fallback. Both are MIT-licensed and download weights from HuggingFace on first use (cached via the existing hf_cache bind mount), unlike U2Net/SAM which need an explicit download script. - config: new BG_REMOVAL_MODEL setting (default "ben2") - tools.py: remove-background-base64 now tries local backends in order (request.model override > BG_REMOVAL_MODEL > ben2/u2net), falling back to rembg's birefnet-general session as a last resort - requirements.gpu.txt / Dockerfile.gpu: add ben2 + transformers deps needed for the new backends, with a build-time smoke test for ben2 - frontend: model dropdown in the Remove Background dialog, threaded through api.js to the new request field Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01Ro4PwQKvSc3CH19LSN21Ht
This commit is contained in:
@@ -46,6 +46,11 @@ class Settings(BaseSettings):
|
||||
# Allow per-edit model override
|
||||
allow_model_override: bool = True
|
||||
|
||||
# Remove Background — preferred local model when request.model="auto"
|
||||
# Options: ben2 (default, best for clean cutouts/hair), birefnet-hr (best
|
||||
# for high-res/print work), u2net (lightweight, smallest download)
|
||||
bg_removal_model: str = "ben2"
|
||||
|
||||
# Local GPU diffusion (AI_PROVIDER=local_gpu)
|
||||
auto_download_models: bool = True # download HF models on first use
|
||||
local_gpu_max_pipelines: int = 2 # max diffusion pipelines kept in GPU memory
|
||||
|
||||
+112
-15
@@ -35,6 +35,7 @@ class InpaintRequest(BaseModel):
|
||||
|
||||
class RemoveBackgroundRequest(BaseModel):
|
||||
image: str # Base64 encoded image
|
||||
model: Optional[str] = "auto" # "auto", "ben2", "birefnet-hr", "u2net", "rembg"
|
||||
|
||||
|
||||
@router.post("/smart-select-base64")
|
||||
@@ -128,36 +129,54 @@ async def inpaint_base64(request: InpaintRequest):
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
class RemoveBackgroundRequestV2(BaseModel):
|
||||
image: str # Base64 encoded image
|
||||
model: Optional[str] = "auto" # "auto", "u2net", "rembg", "birefnet"
|
||||
|
||||
|
||||
@router.post("/remove-background-base64")
|
||||
async def remove_background_base64(request: RemoveBackgroundRequest):
|
||||
"""
|
||||
Remove background from a base64 encoded image.
|
||||
Tries multiple methods: U2Net (direct), rembg with BiRefNet, rembg default.
|
||||
|
||||
request.model selects the backend:
|
||||
- "auto" (default): BG_REMOVAL_MODEL setting first, then falls back
|
||||
through the other local models, then rembg as a last resort.
|
||||
- "ben2" / "birefnet-hr" / "u2net": use only that local model.
|
||||
- "rembg": skip local models, use rembg directly.
|
||||
|
||||
Returns base64 encoded PNG with transparent background.
|
||||
Used by miniPaint frontend.
|
||||
"""
|
||||
try:
|
||||
from app.config import settings
|
||||
|
||||
# Decode base64 image
|
||||
image_bytes = base64.b64decode(request.image)
|
||||
img = Image.open(BytesIO(image_bytes)).convert('RGB')
|
||||
|
||||
local_backends = {
|
||||
"ben2": _remove_background_ben2,
|
||||
"birefnet-hr": _remove_background_birefnet_hr,
|
||||
"u2net": _remove_background_u2net,
|
||||
}
|
||||
|
||||
if request.model in local_backends:
|
||||
order = [request.model]
|
||||
elif request.model == "rembg":
|
||||
order = []
|
||||
else:
|
||||
preferred = settings.bg_removal_model if settings.bg_removal_model in local_backends else "ben2"
|
||||
order = [preferred] + [name for name in ("ben2", "u2net") if name != preferred]
|
||||
|
||||
result_bytes = None
|
||||
method_used = None
|
||||
|
||||
# Try U2Net first (direct implementation, no rembg dependency issues)
|
||||
try:
|
||||
result_bytes = await _remove_background_u2net(img)
|
||||
method_used = "u2net"
|
||||
except Exception as e:
|
||||
print(f"U2Net failed: {e}")
|
||||
for name in order:
|
||||
try:
|
||||
result_bytes = await local_backends[name](img)
|
||||
method_used = name
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"{name} failed: {e}")
|
||||
|
||||
# Fall back to rembg if U2Net failed
|
||||
if result_bytes is None:
|
||||
# rembg is the universal last resort (also reachable directly via model="rembg")
|
||||
if result_bytes is None and request.model in ("auto", "rembg"):
|
||||
try:
|
||||
from rembg import remove, new_session
|
||||
try:
|
||||
@@ -175,7 +194,7 @@ async def remove_background_base64(request: RemoveBackgroundRequest):
|
||||
if result_bytes is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="No background removal method available. Install u2net or rembg."
|
||||
detail="No background removal method available. Install ben2, u2net, or rembg."
|
||||
)
|
||||
|
||||
# Convert result to base64
|
||||
@@ -328,6 +347,84 @@ async def _remove_background_u2net(img: Image.Image) -> bytes:
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
# Global BEN2 model cache
|
||||
_ben2_model = None
|
||||
|
||||
|
||||
async def _remove_background_ben2(img: Image.Image) -> bytes:
|
||||
"""
|
||||
Remove background using BEN2 (Confidence Guided Matting) — clean cutouts,
|
||||
strong on hair/fur edges. MIT licensed. Downloads weights from HF Hub on
|
||||
first use (cached under the hf_cache bind mount).
|
||||
"""
|
||||
global _ben2_model
|
||||
|
||||
if _ben2_model is None:
|
||||
import torch
|
||||
from ben2 import AutoModel as Ben2AutoModel
|
||||
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
print(f"Loading BEN2_Base model on {device} (first run downloads ~170MB from HuggingFace)")
|
||||
_ben2_model = Ben2AutoModel.from_pretrained("PramaLLC/BEN2")
|
||||
_ben2_model.to(device).eval()
|
||||
print("BEN2_Base model loaded")
|
||||
|
||||
result = _ben2_model.inference(img.convert('RGB'), refine_foreground=False)
|
||||
|
||||
buffer = BytesIO()
|
||||
result.save(buffer, format='PNG')
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
# Global BiRefNet-HR model cache
|
||||
_birefnet_hr_model = None
|
||||
_birefnet_hr_device = None
|
||||
|
||||
|
||||
async def _remove_background_birefnet_hr(img: Image.Image) -> bytes:
|
||||
"""
|
||||
Remove background using BiRefNet-HR (2048x2048, MIT licensed) — best for
|
||||
high-resolution / print work. Downloads weights from HF Hub on first use.
|
||||
"""
|
||||
global _birefnet_hr_model, _birefnet_hr_device
|
||||
|
||||
import torch
|
||||
from torchvision import transforms
|
||||
|
||||
if _birefnet_hr_model is None:
|
||||
from transformers import AutoModelForImageSegmentation
|
||||
|
||||
_birefnet_hr_device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
print(f"Loading BiRefNet-HR model on {_birefnet_hr_device} (first run downloads ~900MB from HuggingFace)")
|
||||
_birefnet_hr_model = AutoModelForImageSegmentation.from_pretrained(
|
||||
'zhengpeng7/BiRefNet_HR', trust_remote_code=True
|
||||
)
|
||||
_birefnet_hr_model.to(_birefnet_hr_device).eval()
|
||||
print("BiRefNet-HR model loaded")
|
||||
|
||||
original_size = img.size
|
||||
rgb_img = img.convert('RGB')
|
||||
|
||||
transform = transforms.Compose([
|
||||
transforms.Resize((2048, 2048)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
])
|
||||
input_tensor = transform(rgb_img).unsqueeze(0).to(_birefnet_hr_device)
|
||||
|
||||
with torch.no_grad():
|
||||
preds = _birefnet_hr_model(input_tensor)[-1].sigmoid().cpu()
|
||||
|
||||
mask = transforms.ToPILImage()(preds[0].squeeze()).resize(original_size, Image.Resampling.LANCZOS)
|
||||
|
||||
result = rgb_img.convert('RGBA')
|
||||
result.putalpha(mask)
|
||||
|
||||
buffer = BytesIO()
|
||||
result.save(buffer, format='PNG')
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
@router.post("/remove-background")
|
||||
async def remove_background(
|
||||
project_id: Optional[int] = Form(None),
|
||||
|
||||
@@ -31,3 +31,14 @@ sentencepiece>=0.2.0
|
||||
# Install post-container-start if needed:
|
||||
# pip install xformers --index-url https://download.pytorch.org/whl/cu121
|
||||
# xformers
|
||||
|
||||
# Background removal — BEN2 (default, clean cutouts/hair) + BiRefNet-HR
|
||||
# (high-res/print alternate). Both MIT-licensed. Verified against upstream
|
||||
# source: neither requires torch>=2.5 despite the BiRefNet repo's own
|
||||
# requirements.txt floor — that pin is for its training/eval scripts, not
|
||||
# the inference path used here. Weights download from HuggingFace on first
|
||||
# use (cached via the hf_cache bind mount, same as the diffusion models).
|
||||
ben2 @ git+https://github.com/PramaLLC/BEN2.git
|
||||
timm>=1.0.10
|
||||
einops>=0.6.0
|
||||
kornia>=0.7.0
|
||||
|
||||
Reference in New Issue
Block a user