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:
Claude
2026-06-18 14:20:24 +00:00
parent 781b9e7cdc
commit 38e80f8fb8
9 changed files with 169 additions and 32 deletions
+5
View File
@@ -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
View File
@@ -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),