mirror of
https://github.com/snapotter-hq/SnapOtter.git
synced 2026-08-03 07:46:42 +02:00
feat: add Ultra quality mode with BiRefNet-matting for people photos
Adds a new "Ultra" quality tier for People subject type that uses BiRefNet-matting (ONNX, 928MB) for true alpha matting instead of binary segmentation. Produces per-pixel transparency for hair wisps and fine edges that standard models miss. - Custom rembg session class loads BiRefNet-matting ONNX from GitHub releases - Zero new Python dependencies (reuses existing onnxruntime) - Model pre-downloaded in Docker build alongside existing models - Ultra option only visible when subject is People - Falls back to Best when switching to Products/General
This commit is contained in:
@@ -9,6 +9,39 @@ def emit_progress(percent, stage):
|
||||
print(json.dumps({"progress": percent, "stage": stage}), file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
_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)
|
||||
|
||||
|
||||
def main():
|
||||
input_path = sys.argv[1]
|
||||
output_path = sys.argv[2]
|
||||
@@ -23,8 +56,12 @@ def main():
|
||||
|
||||
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)
|
||||
|
||||
emit_progress(10, "Loading model")
|
||||
|
||||
session = new_session(model, providers=onnx_providers())
|
||||
|
||||
Reference in New Issue
Block a user