From 08378e57f87d0f1ece8a4073b4befd439a0239d9 Mon Sep 17 00:00:00 2001 From: lingh Date: Tue, 30 Jun 2026 23:58:22 +0800 Subject: [PATCH] Add anime-seg segmentation backend (SkyTNT ISNet) Add AnimeSegSegmenter (skytnt/anime-seg ISNet ONNX via onnxruntime, no remote code) and a make_segmenter factory selected by SegmentationSettings.backend ("birefnet" | "anime-seg"). The ONNX output is already 0..1, so it slots into the same soft-mask interface BiRefNetSegmenter uses. On the anime samples anime-seg recovers more and more-coherent hair wisps than BiRefNet (TestImage2 shoulder rescue: added px 1886 -> 3278, largest connected component 147 -> 454), as expected from an anime-trained model. Default backend stays birefnet. Adds onnxruntime to requirements. Co-Authored-By: Claude Opus 4.8 --- bgfilter/pipeline.py | 4 ++-- bgfilter/segmentation.py | 48 ++++++++++++++++++++++++++++++++++++++++ bgfilter/settings.py | 1 + configs/default.yaml | 1 + requirements.txt | 1 + 5 files changed, 53 insertions(+), 2 deletions(-) diff --git a/bgfilter/pipeline.py b/bgfilter/pipeline.py index 2eebf6a..1c30a13 100644 --- a/bgfilter/pipeline.py +++ b/bgfilter/pipeline.py @@ -31,9 +31,9 @@ class MattingPipeline: def _segment(self, rgb: np.ndarray) -> np.ndarray: if self._segmenter is None: - from .segmentation import BiRefNetSegmenter + from .segmentation import make_segmenter - self._segmenter = BiRefNetSegmenter(self.settings.segmentation) + self._segmenter = make_segmenter(self.settings.segmentation) return self._segmenter.mask(rgb) def _predict_alpha(self, rgb: np.ndarray, trimap: np.ndarray, bg_confidence: np.ndarray) -> tuple[np.ndarray, str]: diff --git a/bgfilter/segmentation.py b/bgfilter/segmentation.py index 4328d11..e6269bf 100644 --- a/bgfilter/segmentation.py +++ b/bgfilter/segmentation.py @@ -71,3 +71,51 @@ class BiRefNetSegmenter: ).resize((rgb.shape[1], rgb.shape[0]), Image.Resampling.BILINEAR) pred = np.asarray(pred_img, dtype=np.float32) / 255.0 return np.clip(pred, 0.0, 1.0).astype(np.float32) + + +class AnimeSegSegmenter: + """Anime character segmentation via SkyTNT anime-seg (ISNet, ONNX). + + Anime-trained, run through onnxruntime (no remote code). Same soft-mask + interface as :class:`BiRefNetSegmenter`. The ONNX output is already 0..1. + Weights load from HuggingFace; behind a firewall set ``HF_ENDPOINT``. + """ + + def __init__(self, settings: SegmentationSettings): + try: + import onnxruntime as ort + from huggingface_hub import hf_hub_download + except ModuleNotFoundError as exc: + raise RuntimeError( + "Missing anime-seg dependencies. Install them with: " + "pip install onnxruntime huggingface_hub" + ) from exc + + model_file = hf_hub_download(settings.model_name, "isnetis.onnx") + providers = ["CPUExecutionProvider"] + if settings.device == "cuda" and "CUDAExecutionProvider" in ort.get_available_providers(): + providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] + self.session = ort.InferenceSession(model_file, providers=providers) + self.input_name = self.session.get_inputs()[0].name + self.size = settings.input_size + + def mask(self, rgb: np.ndarray) -> np.ndarray: + image = Image.fromarray(rgb.astype(np.uint8), mode="RGB").resize( + (self.size, self.size), Image.Resampling.BILINEAR + ) + x = (np.asarray(image, dtype=np.float32) / 255.0).transpose(2, 0, 1)[None] + pred = self.session.run(None, {self.input_name: x})[0][0, 0] + mask_img = Image.fromarray( + np.clip(pred * 255.0, 0, 255).astype(np.uint8), mode="L" + ).resize((rgb.shape[1], rgb.shape[0]), Image.Resampling.BILINEAR) + return np.asarray(mask_img, dtype=np.float32) / 255.0 + + +def make_segmenter(settings: SegmentationSettings): + if settings.backend == "birefnet": + return BiRefNetSegmenter(settings) + if settings.backend == "anime-seg": + return AnimeSegSegmenter(settings) + raise RuntimeError( + f"Unknown segmentation backend '{settings.backend}'. Use 'birefnet' or 'anime-seg'." + ) diff --git a/bgfilter/settings.py b/bgfilter/settings.py index 100d4ff..6c27e6e 100644 --- a/bgfilter/settings.py +++ b/bgfilter/settings.py @@ -104,6 +104,7 @@ class ModelSettings: @dataclass(frozen=True) class SegmentationSettings: enabled: bool = True + backend: str = "birefnet" # "birefnet" or "anime-seg" model_name: str = "ZhengPeng7/BiRefNet" device: str = "cuda" input_size: int = 1024 diff --git a/configs/default.yaml b/configs/default.yaml index 3f9a95d..abae22e 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -6,6 +6,7 @@ model: segmentation: enabled: true + backend: birefnet # birefnet | anime-seg (anime-seg needs model_name: skytnt/anime-seg) model_name: ZhengPeng7/BiRefNet device: cuda input_size: 1024 diff --git a/requirements.txt b/requirements.txt index 6fddb5d..b36b1f9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,6 +13,7 @@ huggingface_hub timm einops kornia +onnxruntime tqdm typer rich