mirror of
https://github.com/snapotter-hq/SnapOtter.git
synced 2026-08-03 07:46:42 +02:00
fix: guard against division-by-zero in HR-matting predict and register session in installer
The BiRefNetHRMattingSession.predict normalization crashes when all pixels share the same value (ma == mi). Use a guarded denominator so uniform-alpha inputs produce a zero mask instead of a NaN explosion. Also adds _register_birefnet_hr_matting() to install_feature.py so the HR-matting model can be downloaded during feature installation, matching the existing registration in remove_bg.py.
This commit is contained in:
@@ -316,6 +316,63 @@ def _register_birefnet_matting() -> None:
|
|||||||
sessions_class.append(BiRefNetMattingSession)
|
sessions_class.append(BiRefNetMattingSession)
|
||||||
|
|
||||||
|
|
||||||
|
_hr_matting_registered = False
|
||||||
|
|
||||||
|
|
||||||
|
def _register_birefnet_hr_matting() -> None:
|
||||||
|
"""Register the custom BiRefNet HR-matting ONNX session for 2048x2048 high-res matting.
|
||||||
|
|
||||||
|
Like _register_birefnet_matting(), this model is not built into rembg and
|
||||||
|
must be registered before calling new_session("birefnet-hr-matting").
|
||||||
|
"""
|
||||||
|
global _hr_matting_registered
|
||||||
|
if _hr_matting_registered:
|
||||||
|
return
|
||||||
|
_hr_matting_registered = True
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pooch
|
||||||
|
from PIL import Image
|
||||||
|
from rembg.sessions import sessions_class
|
||||||
|
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 download_rembg_session(model: dict, models_dir: str) -> None:
|
def download_rembg_session(model: dict, models_dir: str) -> None:
|
||||||
"""Download a rembg model by initializing a session."""
|
"""Download a rembg model by initializing a session."""
|
||||||
args = model.get("args", [])
|
args = model.get("args", [])
|
||||||
@@ -331,6 +388,7 @@ def download_rembg_session(model: dict, models_dir: str) -> None:
|
|||||||
|
|
||||||
from rembg import new_session
|
from rembg import new_session
|
||||||
_register_birefnet_matting()
|
_register_birefnet_matting()
|
||||||
|
_register_birefnet_hr_matting()
|
||||||
new_session(model_name)
|
new_session(model_name)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -94,7 +94,8 @@ def _register_hr_matting_session(sessions_class):
|
|||||||
pred = ort_outs[0][:, 0, :, :]
|
pred = ort_outs[0][:, 0, :, :]
|
||||||
ma = np.max(pred)
|
ma = np.max(pred)
|
||||||
mi = np.min(pred)
|
mi = np.min(pred)
|
||||||
pred = (pred - mi) / (ma - mi)
|
denom = ma - mi
|
||||||
|
pred = (pred - mi) / denom if denom > 0 else pred * 0
|
||||||
pred = np.squeeze(pred)
|
pred = np.squeeze(pred)
|
||||||
mask = Image.fromarray((pred * 255).astype("uint8"), mode="L")
|
mask = Image.fromarray((pred * 255).astype("uint8"), mode="L")
|
||||||
mask = mask.resize(img.size, Image.LANCZOS)
|
mask = mask.resize(img.size, Image.LANCZOS)
|
||||||
|
|||||||
Reference in New Issue
Block a user