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:
@@ -13,18 +13,24 @@ import { useToolProcessor } from "@/hooks/use-tool-processor";
|
|||||||
import { useFileStore } from "@/stores/file-store";
|
import { useFileStore } from "@/stores/file-store";
|
||||||
|
|
||||||
type SubjectType = "people" | "products" | "general";
|
type SubjectType = "people" | "products" | "general";
|
||||||
type Quality = "fast" | "balanced" | "best";
|
type Quality = "fast" | "balanced" | "best" | "ultra";
|
||||||
type BackgroundType = "transparent" | "color" | "gradient" | "image";
|
type BackgroundType = "transparent" | "color" | "gradient" | "image";
|
||||||
|
|
||||||
type BgModel =
|
type BgModel =
|
||||||
| "birefnet-general"
|
| "birefnet-general"
|
||||||
| "birefnet-general-lite"
|
| "birefnet-general-lite"
|
||||||
|
| "birefnet-matting"
|
||||||
| "birefnet-portrait"
|
| "birefnet-portrait"
|
||||||
| "bria-rmbg"
|
| "bria-rmbg"
|
||||||
| "u2net";
|
| "u2net";
|
||||||
|
|
||||||
const MODEL_MAP: Record<SubjectType, Record<Quality, BgModel>> = {
|
const MODEL_MAP: Record<SubjectType, Partial<Record<Quality, BgModel>>> = {
|
||||||
people: { fast: "u2net", balanced: "birefnet-portrait", best: "birefnet-portrait" },
|
people: {
|
||||||
|
fast: "u2net",
|
||||||
|
balanced: "birefnet-portrait",
|
||||||
|
best: "birefnet-portrait",
|
||||||
|
ultra: "birefnet-matting",
|
||||||
|
},
|
||||||
products: { fast: "u2net", balanced: "bria-rmbg", best: "birefnet-general" },
|
products: { fast: "u2net", balanced: "bria-rmbg", best: "birefnet-general" },
|
||||||
general: { fast: "u2net", balanced: "birefnet-general-lite", best: "birefnet-general" },
|
general: { fast: "u2net", balanced: "birefnet-general-lite", best: "birefnet-general" },
|
||||||
};
|
};
|
||||||
@@ -35,10 +41,11 @@ const SUBJECT_OPTIONS: { value: SubjectType; label: string; icon: typeof User }[
|
|||||||
{ value: "general", label: "General", icon: ImageIcon },
|
{ value: "general", label: "General", icon: ImageIcon },
|
||||||
];
|
];
|
||||||
|
|
||||||
const QUALITY_OPTIONS: { value: Quality; label: string }[] = [
|
const ALL_QUALITY_OPTIONS: { value: Quality; label: string; peopleOnly?: boolean }[] = [
|
||||||
{ value: "fast", label: "Fast" },
|
{ value: "fast", label: "Fast" },
|
||||||
{ value: "balanced", label: "Balanced" },
|
{ value: "balanced", label: "Balanced" },
|
||||||
{ value: "best", label: "Best" },
|
{ value: "best", label: "Best" },
|
||||||
|
{ value: "ultra", label: "Ultra", peopleOnly: true },
|
||||||
];
|
];
|
||||||
|
|
||||||
const COLOR_PRESETS = [
|
const COLOR_PRESETS = [
|
||||||
@@ -97,8 +104,18 @@ export function RemoveBgControls({ settings, onChange }: RemoveBgControlsProps)
|
|||||||
// Expandable sections
|
// Expandable sections
|
||||||
const [effectsOpen, setEffectsOpen] = useState(false);
|
const [effectsOpen, setEffectsOpen] = useState(false);
|
||||||
|
|
||||||
|
// Filter quality options based on subject (Ultra only for People)
|
||||||
|
const qualityOptions = ALL_QUALITY_OPTIONS.filter(
|
||||||
|
(opt) => !opt.peopleOnly || subject === "people",
|
||||||
|
);
|
||||||
|
|
||||||
|
// If switching away from People while on Ultra, fall back to Best
|
||||||
|
const effectiveQuality = quality === "ultra" && subject !== "people" ? "best" : quality;
|
||||||
|
|
||||||
const model =
|
const model =
|
||||||
isPassport && subject === "people" ? "birefnet-portrait" : MODEL_MAP[subject][quality];
|
isPassport && subject === "people"
|
||||||
|
? "birefnet-portrait"
|
||||||
|
: MODEL_MAP[subject][effectiveQuality] || "birefnet-general";
|
||||||
|
|
||||||
const onChangeRef = useRef(onChange);
|
const onChangeRef = useRef(onChange);
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
@@ -190,14 +207,14 @@ export function RemoveBgControls({ settings, onChange }: RemoveBgControlsProps)
|
|||||||
|
|
||||||
{/* Quality */}
|
{/* Quality */}
|
||||||
<SectionLabel>Quality</SectionLabel>
|
<SectionLabel>Quality</SectionLabel>
|
||||||
<div className="grid grid-cols-3 gap-1.5">
|
<div className={`grid gap-1.5 ${qualityOptions.length > 3 ? "grid-cols-4" : "grid-cols-3"}`}>
|
||||||
{QUALITY_OPTIONS.map((opt) => (
|
{qualityOptions.map((opt) => (
|
||||||
<button
|
<button
|
||||||
key={opt.value}
|
key={opt.value}
|
||||||
type="button"
|
type="button"
|
||||||
onClick={() => setQuality(opt.value)}
|
onClick={() => setQuality(opt.value)}
|
||||||
className={`py-2 px-2 rounded-lg border text-xs font-medium transition-colors ${
|
className={`py-2 px-2 rounded-lg border text-xs font-medium transition-colors ${
|
||||||
quality === opt.value
|
effectiveQuality === opt.value
|
||||||
? "border-primary bg-primary/10 text-primary"
|
? "border-primary bg-primary/10 text-primary"
|
||||||
: "border-border text-muted-foreground hover:border-primary/50"
|
: "border-border text-muted-foreground hover:border-primary/50"
|
||||||
}`}
|
}`}
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ REMBG_MODELS = [
|
|||||||
"birefnet-general-lite",
|
"birefnet-general-lite",
|
||||||
"birefnet-portrait",
|
"birefnet-portrait",
|
||||||
"birefnet-general",
|
"birefnet-general",
|
||||||
|
"birefnet-matting",
|
||||||
]
|
]
|
||||||
|
|
||||||
# PaddleOCR language codes (not ISO). German/French/Spanish use "latin" model.
|
# PaddleOCR language codes (not ISO). German/French/Spanish use "latin" model.
|
||||||
@@ -34,11 +35,40 @@ REMBG_MODELS = [
|
|||||||
PADDLEOCR_LANGUAGES = ["en", "ch", "japan", "korean", "latin"]
|
PADDLEOCR_LANGUAGES = ["en", "ch", "japan", "korean", "latin"]
|
||||||
|
|
||||||
|
|
||||||
|
def _register_birefnet_matting():
|
||||||
|
"""Register BiRefNet-matting ONNX session for Ultra quality mode."""
|
||||||
|
import os
|
||||||
|
import pooch
|
||||||
|
from rembg.sessions import sessions_class
|
||||||
|
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 download_rembg_models():
|
def download_rembg_models():
|
||||||
"""Download all rembg ONNX models."""
|
"""Download all rembg ONNX models."""
|
||||||
print("=== Downloading rembg models ===")
|
print("=== Downloading rembg models ===")
|
||||||
from rembg import new_session
|
from rembg import new_session
|
||||||
|
|
||||||
|
_register_birefnet_matting()
|
||||||
|
|
||||||
for model in REMBG_MODELS:
|
for model in REMBG_MODELS:
|
||||||
print(f" Downloading {model}...")
|
print(f" Downloading {model}...")
|
||||||
new_session(model)
|
new_session(model)
|
||||||
|
|||||||
@@ -9,6 +9,39 @@ def emit_progress(percent, stage):
|
|||||||
print(json.dumps({"progress": percent, "stage": stage}), file=sys.stderr, flush=True)
|
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():
|
def main():
|
||||||
input_path = sys.argv[1]
|
input_path = sys.argv[1]
|
||||||
output_path = sys.argv[2]
|
output_path = sys.argv[2]
|
||||||
@@ -23,8 +56,12 @@ def main():
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from rembg import remove, new_session
|
from rembg import remove, new_session
|
||||||
|
from rembg.sessions import sessions_class
|
||||||
from gpu import onnx_providers
|
from gpu import onnx_providers
|
||||||
|
|
||||||
|
# Register BiRefNet-matting (Ultra quality) if not already present
|
||||||
|
_register_matting_session(sessions_class)
|
||||||
|
|
||||||
emit_progress(10, "Loading model")
|
emit_progress(10, "Loading model")
|
||||||
|
|
||||||
session = new_session(model, providers=onnx_providers())
|
session = new_session(model, providers=onnx_providers())
|
||||||
|
|||||||
@@ -108,6 +108,33 @@ test.describe("Remove Background tool", () => {
|
|||||||
await expect(page.locator("section[aria-label='Image area'] img").first()).toBeVisible();
|
await expect(page.locator("section[aria-label='Image area'] img").first()).toBeVisible();
|
||||||
});
|
});
|
||||||
|
|
||||||
|
test("Ultra quality visible for People, hidden for Products", async ({ loggedInPage: page }) => {
|
||||||
|
await page.goto("/remove-background");
|
||||||
|
|
||||||
|
// People is default - Ultra should be visible
|
||||||
|
await expect(page.getByRole("button", { name: "Ultra" })).toBeVisible();
|
||||||
|
|
||||||
|
// Switch to Products - Ultra should disappear
|
||||||
|
await page.getByText("Products").click();
|
||||||
|
await expect(page.getByRole("button", { name: "Ultra" })).not.toBeVisible();
|
||||||
|
|
||||||
|
// Switch back to People - Ultra returns
|
||||||
|
await page.getByText("People").click();
|
||||||
|
await expect(page.getByRole("button", { name: "Ultra" })).toBeVisible();
|
||||||
|
});
|
||||||
|
|
||||||
|
test("Ultra quality processes JPG portrait", async ({ loggedInPage: page }) => {
|
||||||
|
await page.goto("/remove-background");
|
||||||
|
await uploadFile(page, fixturePath("test-portrait.jpg"));
|
||||||
|
|
||||||
|
// Select Ultra quality
|
||||||
|
await page.getByRole("button", { name: "Ultra" }).click();
|
||||||
|
|
||||||
|
await removeBgAndWait(page);
|
||||||
|
await expect(page.locator("section[aria-label='Image area'] img").first()).toBeVisible();
|
||||||
|
await expect(page.locator("text=Background removal failed")).not.toBeVisible();
|
||||||
|
});
|
||||||
|
|
||||||
test("HEIC portrait - processes without error", async ({ loggedInPage: page }) => {
|
test("HEIC portrait - processes without error", async ({ loggedInPage: page }) => {
|
||||||
await page.goto("/remove-background");
|
await page.goto("/remove-background");
|
||||||
await uploadFile(page, fixturePath("test-portrait.heic"));
|
await uploadFile(page, fixturePath("test-portrait.heic"));
|
||||||
|
|||||||
Reference in New Issue
Block a user