feat: initial project setup

This commit is contained in:
2026-10-08 09:46:47 +08:00
commit 5833a303bc
65 changed files with 11628 additions and 0 deletions
+14
View File
@@ -0,0 +1,14 @@
# -*- coding: utf-8 -*-
"""BiRefNet WebUI —— 基于 ComfyUI 节点 comfyui_birefnet_ll 的网页版抠图工具。
模块划分:
compat.py 定位 ComfyUI 节点目录、注入 folder_paths 垫片、导入模型类
fsutil.py 与推理无关的小工具(图片后缀、文件名清洗、最终尺寸换算)
imageops.py 图像预处理与后处理(纯 numpy / PIL / torch)
batch.py 目录批量抠图:路径校验、目录扫描、输出命名
engine.py 模型扫描、加载缓存与推理
server.py HTTP 服务、REST API、任务队列
"""
__version__ = "1.0.0"
APP_NAME = "BiRefNet WebUI"
+222
View File
@@ -0,0 +1,222 @@
# -*- coding: utf-8 -*-
"""目录批量抠图:路径校验 / 目录扫描 / 输出文件命名。
三条硬约束(对应需求):
1. **校验阶段只读不写**:输入目录、输出目录都必须**已经存在**,本模块
绝不 ``mkdir``。路径写错时直接报错,而不是顺手把目录树建出来 ——
这样一次误输入(或网页上的恶意路径)不会在磁盘上留下任何东西。
2. **命名固定**:``RMBG_<原主干>_<Unix 时间戳>.<后缀>``,原主干超过
:data:`STEM_MAXLEN` 个字符时截断。主干统一走
:func:`birefnet_web.imageops.safe_stem` 清洗,杜绝 ``..`` 之类的穿越。
3. **后缀自适应**:透明输出需要 alpha 通道,``.jpg`` / ``.bmp`` 装不下,
此时退化为 ``.png``(强行按原后缀保存会被 PIL 静默压成白底图)。
本模块不依赖 torch,可以脱离推理环境单独测试。
"""
from __future__ import annotations
import os
from pathlib import Path
from typing import Iterable, List, Optional, Sequence
from .fsutil import ALPHA_SUFFIXES, IMAGE_SUFFIXES, safe_stem
#: 输出文件名统一前缀(扫描输入目录时也用来自我过滤,避免把产物再抠一遍)
BATCH_PREFIX = "RMBG_"
#: 原文件名主干的最大长度(超过即截断)
STEM_MAXLEN = 20
#: 重名避让时最多尝试的序号
_MAX_DEDUP = 999
class BatchPathError(ValueError):
"""路径不合规:不存在 / 不是目录 / 不可写 / 命中系统目录黑名单。"""
# --------------------------------------------------------------------------- #
# 路径解析与校验
# --------------------------------------------------------------------------- #
def _critical_dirs() -> List[Path]:
"""返回当前系统上「不该往里写文件」的目录(仅用于提示级别的高危拦截)。"""
roots: List[Path] = []
names = ("SystemRoot", "windir", "ProgramFiles", "ProgramFiles(x86)", "ProgramData")
for name in names:
raw = os.environ.get(name)
if not raw:
continue
p = Path(raw)
if p.is_absolute():
roots.append(p)
if os.name != "nt":
roots.extend(Path(p) for p in ("/bin", "/sbin", "/etc", "/usr", "/boot", "/dev", "/proc", "/sys"))
return roots
def resolve_user_path(raw: object, base: Path) -> Path:
"""把用户在界面上敲的一行路径规范化成绝对路径。
Args:
raw: 原始输入(可能带首尾空格、成对引号、``~`` 或 ``%VAR%``)。
base: 相对路径的基准目录(本项目的 ``PROJECT_ROOT``)。
Returns:
规范化后的绝对路径(**不校验存在性**,也不做任何创建)。
Raises:
BatchPathError: 输入为空。
"""
text = str(raw or "").strip().strip('"').strip("'").strip()
if not text:
raise BatchPathError("路径不能为空")
text = os.path.expandvars(os.path.expanduser(text))
path = Path(text)
if not path.is_absolute():
path = base / path
return Path(os.path.normpath(str(path)))
def check_input_dir(raw: object, base: Path) -> Path:
"""校验「图片所在目录」,必须已存在且是目录。
Raises:
BatchPathError: 路径为空 / 不存在 / 不是目录。
"""
path = resolve_user_path(raw, base)
if not path.exists():
raise BatchPathError(f"输入目录不存在:{path}(本工具不会自动创建目录,请先确认路径)")
if not path.is_dir():
raise BatchPathError(f"输入路径不是目录:{path}")
return path
def check_output_dir(raw: object, base: Path, *, allow_system_dir: bool = False) -> Path:
"""校验「结果输出目录」,必须已存在、是目录且可写。
Args:
raw: 用户输入或默认输出目录。
base: 相对路径基准。
allow_system_dir: 为 True 时跳过系统目录黑名单(测试用)。
Raises:
BatchPathError: 路径为空 / 不存在 / 不是目录 / 命中黑名单 / 不可写。
"""
path = resolve_user_path(raw, base)
if not path.exists():
raise BatchPathError(f"输出目录不存在:{path}(本工具不会自动创建目录,请先手动建好)")
if not path.is_dir():
raise BatchPathError(f"输出路径不是目录:{path}")
if not allow_system_dir:
if path == Path(path.anchor): # C:\ 或 /
raise BatchPathError(f"拒绝把结果写进磁盘根目录:{path}")
for critical in _critical_dirs():
if path == critical or critical in path.parents:
raise BatchPathError(f"拒绝把结果写进系统目录:{path}")
if not os.access(path, os.W_OK):
raise BatchPathError(f"输出目录不可写(权限不足):{path}")
return path
# --------------------------------------------------------------------------- #
# 目录扫描
# --------------------------------------------------------------------------- #
def scan_images(
root: Path,
*,
recursive: bool = False,
skip_dirs: Sequence[Path] = (),
) -> List[Path]:
"""扫描目录下的图片文件。
Args:
root: 已校验过的图片目录。
recursive: 是否包含子目录。
skip_dirs: 需要跳过的目录(通常是输出目录,避免读到自己的产物)。
Returns:
排序后的图片路径列表(先按目录、再按文件名,保证同一批次的处理顺序稳定)。
"""
skipped = [Path(d).resolve() for d in skip_dirs]
walker: Iterable[Path] = root.rglob("*") if recursive else root.glob("*")
found: List[Path] = []
for path in walker:
try:
if not path.is_file():
continue
except OSError: # 权限/坏链接
continue
if path.suffix.lower() not in IMAGE_SUFFIXES:
continue
# 本工具自己的产物(RMBG_ 前缀)不参与批量,避免反复抠同一张图
if path.name.startswith(BATCH_PREFIX):
continue
if skipped:
parent = path.parent.resolve()
if any(parent == s or s in parent.parents or parent in s.parents for s in skipped):
continue
found.append(path)
return sorted(found, key=lambda p: (str(p.parent).lower(), p.name.lower()))
# --------------------------------------------------------------------------- #
# 输出命名
# --------------------------------------------------------------------------- #
def choose_suffix(orig_suffix: str, background: str) -> str:
"""决定输出文件后缀:透明模式下如果原后缀装不下 alpha,就退化为 ``.png``。
Args:
orig_suffix: 原文件后缀(含点,如 ``.jpg``)。
background: 背景模式,``transparent`` 或其它(纯色)。
Returns:
合法的输出后缀字符串。
"""
suffix = (orig_suffix or "").lower()
if background == "transparent":
return suffix if suffix in ALPHA_SUFFIXES else ".png"
return suffix if suffix in IMAGE_SUFFIXES else ".png"
def build_output_name(source: Path, timestamp: int, background: str = "transparent") -> str:
"""按 ``RMBG_<主干(≤20)>_<Unix 时间戳>.<后缀>`` 组装输出文件名。
Args:
source: 原始图片路径。
timestamp: Unix 时间戳(秒),由调用方在**处理时**取,便于追溯。
background: 背景模式,决定后缀是否允许退化为 ``.png``。
Returns:
输出文件名(不含目录)。
"""
stem = safe_stem(source.name, fallback="image", maxlen=STEM_MAXLEN)
return f"{BATCH_PREFIX}{stem}_{int(timestamp)}{choose_suffix(source.suffix, background)}"
def unique_path(directory: Path, name: str) -> Path:
"""在同一秒内出现同名输出时,用 ``_1`` / ``_2`` 递增避让,绝不覆盖已有文件。
Args:
directory: 输出目录(已存在)。
name: 目标文件名。
Returns:
可安全写入的完整路径。
Raises:
BatchPathError: 避让序号耗尽(同名文件超过 :data:`_MAX_DEDUP` 个)。
"""
target = directory / name
if not target.exists():
return target
for i in range(1, _MAX_DEDUP + 1):
candidate = directory / f"{target.stem}_{i}{target.suffix}"
if not candidate.exists():
return candidate
raise BatchPathError(f"输出目录中同名文件过多,请检查:{name}")
+262
View File
@@ -0,0 +1,262 @@
# -*- coding: utf-8 -*-
"""模型代码定位与 ComfyUI 解耦层。
本项目把原来 ComfyUI 节点 ``comfyui_birefnet_ll`` 的模型代码**内联**到了
``vendor/comfyui_birefnet_ll/``,因此默认情况下完全自包含,不再依赖外部 ComfyUI 安装。
定位优先级(``find_node_dir``):
1. 命令行 ``--node-dir``
2. 环境变量 ``BIREFNET_NODE_DIR``
3. **项目内 ``vendor/comfyui_birefnet_ll``(默认,开箱即用)**
4. ``config.json`` 中的 ``node_dir`` 字段(兼容旧配置)
5. 自动扫描本机 ComfyUI 安装(便于跟随上游节点升级)
模型包在 import 阶段依赖 ComfyUI 的 ``folder_paths``,但只用于「按文件名查权重路径」
(即骨干网络预训练权重)。本项目始终以 ``bb_pretrained=False`` 构建骨干网络,不会真正
读取这类权重,因此这里注入一个最小实现的 ``folder_paths`` 垫片即可脱离 ComfyUI 运行。
"""
from __future__ import annotations
import importlib
import os
import sys
import types
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Tuple
# --------------------------------------------------------------------------- #
# 路径常量
# --------------------------------------------------------------------------- #
PROJECT_ROOT = Path(__file__).resolve().parent.parent
#: 项目内自带的模型代码(默认来源,随项目一起分发)
VENDOR_DIR = PROJECT_ROOT / "vendor" / "comfyui_birefnet_ll"
#: 项目内自带的权重目录(默认来源)
LOCAL_MODEL_DIR = PROJECT_ROOT / "models"
#: 环境变量:手动指定模型代码目录
NODE_DIR_ENV = "BIREFNET_NODE_DIR"
#: 节点仓库名,用于兜底扫描外部 ComfyUI 安装
NODE_REPO_NAME = "comfyui_birefnet_ll"
#: 旧版架构权重文件名(与 birefnetNode.py 保持一致)
OLD_MODEL_FILES: Tuple[str, ...] = ("BiRefNet-DIS_ep580.pth", "BiRefNet-ep480.pth")
#: 判定「是有效模型代码目录」的标志文件
_MARKER = ("birefnet", "models", "birefnet.py")
_BUNDLED_PYTHON = PROJECT_ROOT / "python" / "python.exe"
_INSTALL_HINT = (
"缺少运行依赖。本项目自带独立 Python 运行时,请使用:\n"
f' "{_BUNDLED_PYTHON}" Webui.py\n'
"或双击 run.bat。\n"
"若想用自己的解释器,请先安装依赖:pip install -r requirements.txt"
)
# --------------------------------------------------------------------------- #
# 模型代码目录定位
# --------------------------------------------------------------------------- #
def _is_code_dir(path: Path) -> bool:
"""目录内是否含 ``birefnet/models/birefnet.py``。"""
try:
return path.joinpath(*_MARKER).is_file()
except OSError:
return False
def _scan_external_node_dirs() -> List[Path]:
"""扫描本机常见的 ComfyUI 安装位置(仅在需要跟随上游升级时才会命中)。"""
found: List[Path] = []
home = Path.home()
roots: List[Path] = [home]
for letter in "CDEFGHIJ":
roots.append(Path(f"{letter}:/"))
patterns = (
"ComfyUI/custom_nodes/" + NODE_REPO_NAME,
"ComfyUI*/ComfyUI/custom_nodes/" + NODE_REPO_NAME,
"ComfyUI*/custom_nodes/" + NODE_REPO_NAME,
"ComfyUI*/ComfyUI/custom_nodes/ComfyUI_BiRefNet_ll",
"ComfyUI*/custom_nodes/ComfyUI_BiRefNet_ll",
"*/*/ComfyUI/custom_nodes/" + NODE_REPO_NAME,
)
for root in roots:
if not root.exists():
continue
for pat in patterns:
try:
found.extend(sorted(root.glob(pat)))
except (OSError, ValueError):
continue
return found
def _dedup(paths: Iterable[Path]) -> List[Path]:
seen, out = set(), []
for p in paths:
try:
key = str(p.resolve()).lower() if p.exists() else str(p).lower()
except OSError:
key = str(p).lower()
if key in seen:
continue
seen.add(key)
out.append(p)
return out
def candidate_code_dirs(config_hint: Optional[str] = None) -> List[Path]:
"""按优先级生成候选模型代码目录。"""
cands: List[Path] = []
env = os.environ.get(NODE_DIR_ENV)
if env:
cands.append(Path(env))
# 项目自带(默认命中)
cands.append(VENDOR_DIR)
# 旧配置中记录的路径(兼容早期版本写死的 ComfyUI 节点路径)
if config_hint:
cands.append(Path(config_hint))
# 外部 ComfyUI 安装(兜底)
cands.append(PROJECT_ROOT.parent / NODE_REPO_NAME)
cands.extend(_scan_external_node_dirs())
return _dedup(cands)
def find_node_dir(explicit: Optional[str] = None, config_hint: Optional[str] = None) -> Path:
"""定位模型代码目录(需含 ``birefnet/models/birefnet.py``)。
正常情况下会直接命中项目内的 ``vendor/comfyui_birefnet_ll``。
"""
cands: List[Path] = []
if explicit:
cands.append(Path(explicit))
cands.extend(candidate_code_dirs(config_hint))
for cand in cands:
if _is_code_dir(cand):
try:
return cand.resolve()
except OSError:
return cand
raise FileNotFoundError(
"未找到 BiRefNet 模型代码目录(应包含 birefnet/models/birefnet.py)。\n"
f"默认位置:{VENDOR_DIR}\n"
"若该目录缺失或被误删,可从 ComfyUI 节点 comfyui_birefnet_ll 重新拷贝,或用以下方式指定:\n"
" 1) 启动参数 --node-dir <路径>\n"
f" 2) 环境变量 {NODE_DIR_ENV}=<路径>\n"
" 3) config.json 中的 node_dir 字段\n"
"已尝试的候选目录:\n - " + "\n - ".join(str(c) for c in cands[:8])
)
def default_model_dirs(code_dir: Optional[Path] = None) -> List[Path]:
"""返回默认权重搜索目录:项目内 ``models/`` 优先,其次外部目录。"""
dirs: List[Path] = [LOCAL_MODEL_DIR]
# 若用户显式指向了外部 ComfyUI 节点,则顺带搜索对应的 models/BiRefNet
if code_dir is not None:
p = Path(code_dir)
try:
if p.resolve() != VENDOR_DIR.resolve():
dirs.append(p.parent.parent / "models" / "BiRefNet")
except OSError:
pass
return _dedup(dirs)
# --------------------------------------------------------------------------- #
# folder_paths 垫片
# --------------------------------------------------------------------------- #
class _FolderPathsShim(types.ModuleType):
"""仅实现 birefnet/config.py 用到的接口,一律返回 None(不查权重)。"""
def __init__(self, model_dirs: Iterable[Path]) -> None:
super().__init__("folder_paths")
self.models_dir = str(next(iter(model_dirs), Path("models")))
self.folder_names_and_paths: Dict[str, Tuple[List[str], set]] = {
"birefnet": ([self.models_dir], {".pt", ".pth", ".safetensors", ".ckpt"})
}
self.supported_pt_extensions = {".pt", ".pth", ".safetensors", ".ckpt"}
def get_folder_paths(self, folder_name: str) -> List[str]:
return list(self.folder_names_and_paths.get(folder_name, ([], None))[0])
def get_full_path(self, folder_name: str, filename: str) -> Optional[str]:
return None
def get_filename_list(self, folder_name: str) -> List[str]:
return []
def add_model_folder_path(self, folder_name: str, full_folder_path: str) -> None:
if folder_name not in self.folder_names_and_paths:
self.folder_names_and_paths[folder_name] = ([], set())
self.folder_names_and_paths[folder_name][0].append(str(full_folder_path))
def filter_files_extensions(self, files, extensions): # pragma: no cover
return [f for f in files if os.path.splitext(f)[1] in extensions]
def ensure_folder_paths(model_dirs: Iterable[Path]) -> bool:
"""确保 ``folder_paths`` 可用。返回 True 表示用了真实实现(即运行在 ComfyUI 内)。"""
if "folder_paths" in sys.modules:
return not isinstance(sys.modules["folder_paths"], _FolderPathsShim)
try: # 若真的在 ComfyUI 环境内运行,优先用真实实现
importlib.import_module("folder_paths")
return True
except Exception:
pass
sys.modules["folder_paths"] = _FolderPathsShim(model_dirs)
return False
# --------------------------------------------------------------------------- #
# 模型类导入
# --------------------------------------------------------------------------- #
_cached: Dict[str, object] = {}
def load_model_classes(node_dir: Path, model_dirs: Iterable[Path]) -> Dict[str, object]:
"""导入并返回 ``{"BiRefNet":..., "OldBiRefNet":..., "check_state_dict":..., "shim": bool}``。
结果会被缓存,重复调用不会重复导入。
"""
if _cached:
return _cached
node_dir = Path(node_dir).resolve()
shim = ensure_folder_paths(model_dirs)
p = str(node_dir)
if p not in sys.path:
sys.path.insert(0, p)
try:
from birefnet.models.birefnet import BiRefNet
from birefnet.utils import check_state_dict
except ImportError as exc: # pragma: no cover - 环境问题
raise ImportError(f"导入 birefnet 模型包失败({node_dir}):{exc}\n{_INSTALL_HINT}") from exc
try:
from birefnet_old.models.birefnet import BiRefNet as OldBiRefNet
except Exception: # 旧包缺失不影响新模型使用
OldBiRefNet = None # type: ignore[assignment]
_cached.update(
BiRefNet=BiRefNet,
OldBiRefNet=OldBiRefNet,
check_state_dict=check_state_dict,
node_dir=node_dir,
used_folder_paths_shim=shim,
)
return _cached
+499
View File
@@ -0,0 +1,499 @@
# -*- coding: utf-8 -*-
"""模型注册表与推理引擎。
设计要点:
* 模型代码直接复用 ComfyUI 节点(见 compat.py),不复制、不魔改;
* 模型按 (路径, mtime, 设备, 精度, 架构) 做 LRU 缓存,默认只驻留 1 个,避免显存爆炸;
* float16 / bfloat16 推理若出现 NaN/Inf 会自动回退 float32 重算一次并记录告警。
"""
from __future__ import annotations
import gc
import os
import threading
import time
from collections import OrderedDict
from dataclasses import dataclass, field
from pathlib import Path
from typing import Callable, Dict, List, Optional, Tuple
import numpy as np
import torch
from PIL import Image
from . import compat, imageops
ProgressFn = Optional[Callable[[str, float], None]]
#: 不是抠图模型,而是骨干网络权重,扫描时跳过
_BACKBONE_PREFIXES = ("swin_", "pvt_v2_")
#: 支持加载的权重后缀
_MODEL_SUFFIXES = (".safetensors", ".pth", ".pt", ".ckpt")
#: 精度选项 -> torch dtype;auto 由设备决定
DTYPES: Dict[str, Optional[torch.dtype]] = {
"auto": None,
"float32": torch.float32,
"float16": torch.float16,
"bfloat16": torch.bfloat16,
}
#: 前景精修超过该像素量时自动跳过(避免内存与耗时失控)
_REFINE_PIXEL_LIMIT = 40_000_000
@dataclass
class ModelInfo:
"""一个可用的模型权重文件。"""
key: str
name: str
file: str
path: str
size_mb: float
mtime: float
arch: str # "v1" | "old"
bb_index: int
backbone: str
directory: str
def to_dict(self) -> dict:
return {
"key": self.key,
"name": self.name,
"file": self.file,
"path": self.path,
"size_mb": round(self.size_mb, 1),
"arch": self.arch,
"bb_index": self.bb_index,
"backbone": self.backbone,
"directory": self.directory,
}
def guess_bb_index(filename: str) -> int:
"""根据文件名猜骨干网络:lite 系列用 swin_v1_t(3),其余用 swin_v1_l(6)。"""
return 3 if "lite" in filename.lower() else 6
def guess_arch(filename: str) -> str:
return "old" if os.path.basename(filename) in compat.OLD_MODEL_FILES else "v1"
class ModelRegistry:
"""扫描并维护可用模型列表。"""
def __init__(self, model_dirs: List[Path]) -> None:
self.model_dirs = [Path(d) for d in model_dirs]
self._models: "OrderedDict[str, ModelInfo]" = OrderedDict()
self._lock = threading.Lock()
def scan(self) -> List[ModelInfo]:
found: "OrderedDict[str, ModelInfo]" = OrderedDict()
for idx, d in enumerate(self.model_dirs):
if not d or not d.is_dir():
continue
try:
entries = sorted(d.iterdir(), key=lambda p: p.name.lower())
except OSError:
continue
for entry in entries:
if not entry.is_file():
continue
if entry.suffix.lower() not in _MODEL_SUFFIXES:
continue
if entry.name.lower().startswith(_BACKBONE_PREFIXES):
continue
try:
stat = entry.stat()
except OSError:
continue
stem = entry.stem
bb_index = guess_bb_index(stem)
bb = {3: "swin_v1_t", 6: "swin_v1_l"}[bb_index]
key = f"{stem}@{idx}" if stem in found else stem
if key in found: # 极少数重名情况
key = f"{stem}@{idx}"
found[key] = ModelInfo(
key=key,
name=stem,
file=entry.name,
path=str(entry.resolve()),
size_mb=stat.st_size / (1024 * 1024),
mtime=stat.st_mtime,
arch=guess_arch(entry.name),
bb_index=bb_index,
backbone=bb,
directory=str(d),
)
with self._lock:
self._models = found
return list(found.values())
def list(self) -> List[ModelInfo]:
with self._lock:
if not self._models:
return self.scan()
return list(self._models.values())
def get(self, key: str) -> ModelInfo:
for info in self.list():
if key in (info.key, info.name, info.file, info.path):
return info
raise KeyError(f"未找到模型 {key!r},可用模型:{[m.key for m in self.list()]}")
class InferenceEngine:
"""串行执行抠图推理(GPU 上同一时刻只跑一个模型)。"""
def __init__(self, registry: ModelRegistry, node_dir: Path, max_cached: int = 1) -> None:
self.registry = registry
self.node_dir = Path(node_dir)
self.max_cached = max(1, int(max_cached))
self._cache: "OrderedDict[tuple, Tuple[object, str, torch.dtype]]" = OrderedDict()
self._lock = threading.RLock()
self._classes: Optional[dict] = None
self.last_warning: Optional[str] = None
# ------------------------------------------------------------------ #
# 模型加载
# ------------------------------------------------------------------ #
def _model_classes(self) -> dict:
if self._classes is None:
self._classes = compat.load_model_classes(self.node_dir, self.registry.model_dirs)
return self._classes
@staticmethod
def resolve_device(device: str) -> str:
device = (device or "auto").strip().lower()
if device in ("auto", "", "gpu"):
return "cuda" if torch.cuda.is_available() else "cpu"
if device.startswith("cuda"):
if not torch.cuda.is_available():
raise RuntimeError("请求使用 CUDA,但当前 torch 未检测到可用 GPU")
return device
return "cpu"
@staticmethod
def resolve_dtype(dtype: str, device: str) -> torch.dtype:
dt = DTYPES.get((dtype or "auto").lower(), None)
if dt is not None:
return dt
# auto:CUDA 下用 fp32 权重 + fp16 autocast(官方推荐),CPU 下 fp32
return torch.float32
@staticmethod
def _autocast_dtype(dtype: str, device: str) -> Optional[torch.dtype]:
if not device.startswith("cuda"):
return None
key = (dtype or "auto").lower()
if key in ("float16", "auto"):
return torch.float16
if key == "bfloat16":
return torch.bfloat16
return None
def _read_state_dict(self, info: ModelInfo) -> dict:
if info.path.lower().endswith(".safetensors"):
import safetensors.torch
return safetensors.torch.load_file(info.path, device="cpu")
try:
sd = torch.load(info.path, map_location="cpu", weights_only=True)
except Exception:
sd = torch.load(info.path, map_location="cpu", weights_only=False)
check_state_dict = self._model_classes()["check_state_dict"]
return check_state_dict(sd)
def _build_model(self, info: ModelInfo, state_dict: dict, dtype: torch.dtype):
"""构建网络并载入权重;骨干索引猜错时自动换一个重试。"""
classes = self._model_classes()
BiRefNet = classes["BiRefNet"]
OldBiRefNet = classes["OldBiRefNet"]
if info.arch == "old":
if OldBiRefNet is None:
raise RuntimeError("该权重属于旧版架构,但节点内缺少 birefnet_old 包")
candidates = [-1]
else:
candidates = [info.bb_index] + [i for i in (6, 3) if i != info.bb_index]
last_err: Optional[Exception] = None
for bb_index in candidates:
try:
if bb_index < 0:
model = OldBiRefNet(bb_pretrained=False)
version = "old"
else:
model = BiRefNet(bb_pretrained=False, bb_index=bb_index)
version = "v1"
if dtype != torch.float32:
model = model.to(dtype=dtype)
model.load_state_dict(state_dict)
if bb_index >= 0 and bb_index != info.bb_index:
self.last_warning = (
f"模型 {info.name} 的骨干索引自动修正为 {bb_index}(文件名推断为 {info.bb_index})"
)
return model, version, bb_index
except Exception as exc: # 尺寸不匹配等
last_err = exc
continue
raise RuntimeError(f"加载模型 {info.name} 失败:{last_err}") from last_err
def get_model(self, key: str, device: str = "auto", dtype: str = "auto", arch: str = "auto"):
info = self.registry.get(key)
if arch in ("v1", "old"):
info = ModelInfo(**{**info.__dict__, "arch": arch})
device = self.resolve_device(device)
torch_dtype = self.resolve_dtype(dtype, device)
cache_key = (info.path, info.mtime, device, str(torch_dtype), info.arch)
with self._lock:
if cache_key in self._cache:
self._cache.move_to_end(cache_key)
model, version, _ = self._cache[cache_key]
return model, version, info
state_dict = self._read_state_dict(info)
model, version, _ = self._build_model(info, state_dict, torch_dtype)
del state_dict
model = model.to(device)
model.eval()
for p in model.parameters():
p.requires_grad_(False)
self._cache[cache_key] = (model, version, torch_dtype)
self._cache.move_to_end(cache_key)
self._evict_locked()
return model, version, info
def _evict_locked(self) -> None:
while len(self._cache) > self.max_cached:
_, (model, _, _) = self._cache.popitem(last=False)
del model
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def unload(self) -> None:
with self._lock:
self._cache.clear()
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
# ------------------------------------------------------------------ #
# 推理
# ------------------------------------------------------------------ #
def remove_background(
self,
image: Image.Image,
options: Dict[str, object],
progress: ProgressFn = None,
) -> Dict[str, object]:
"""对单张图片执行抠图。
options 支持(均有默认值):
model / device / dtype / arch
resolution_mode: square | longest | custom
width / height / longest_side
upscale_method / mask_threshold
refine_foreground / blur_size / blur_size_two
background: transparent | color
bg_color: "#ffffff"
output_mask: bool(默认 False:不额外产出遮罩文件)
final_longest_side: int (0 = 保持原图尺寸;>0 = 等比缩放到该最长边)
"""
t0 = time.perf_counter()
self.last_warning = None
def report(stage: str, frac: float) -> None:
if progress:
try:
progress(stage, max(0.0, min(1.0, frac)))
except Exception:
pass
model_key = str(options.get("model") or "")
if not model_key:
raise ValueError("未指定模型")
device = str(options.get("device") or "auto")
dtype = str(options.get("dtype") or "auto")
arch = str(options.get("arch") or "auto")
report("加载模型", 0.05)
model, version, info = self.get_model(model_key, device, dtype, arch)
model_dtype = next(model.parameters()).dtype
model_device = next(model.parameters()).device
report("预处理", 0.2)
rgb = imageops.pil_to_rgb_array(image)
src_h, src_w = rgb.shape[:2]
in_h, in_w = imageops.build_input_size(
src_h,
src_w,
mode=str(options.get("resolution_mode") or "square"),
width=int(options.get("width") or 1024),
height=int(options.get("height") or 1024),
longest_side=int(options.get("longest_side") or 1024),
)
upscale_method = str(options.get("upscale_method") or "bilinear")
tensor = imageops.preprocess(rgb, (in_h, in_w), upscale_method)
report("推理", 0.35)
x = tensor.to(model_device)
if model_dtype != torch.float32:
x = x.to(model_dtype)
autocast_dtype = self._autocast_dtype(dtype, device)
alpha = self._forward(model, x, model_device, model_dtype, autocast_dtype)
del x
# 精度兜底:半精度偶发 NaN 时用 fp32 重算一次
if not bool(torch.isfinite(alpha).all()):
self.last_warning = "半精度推理出现 NaN/Inf,已自动回退 float32 重算"
self.unload()
model, version, info = self.get_model(model_key, device, "float32", arch)
model_device = next(model.parameters()).device
x = tensor.to(model_device)
alpha = self._forward(model, x, model_device, torch.float32, None)
del x
del tensor
report("后处理", 0.75)
alpha = imageops.upscale_mask(alpha, src_h, src_w, upscale_method)
alpha = imageops.filter_mask(alpha, float(options.get("mask_threshold") or 0.0))
alpha_np = alpha.squeeze(0).squeeze(0).to(torch.float32).cpu().numpy()
del alpha
alpha_np = np.clip(alpha_np, 0.0, 1.0)
want_mask = bool(options.get("output_mask", False))
background = str(options.get("background") or "transparent")
rgb_out = rgb
refine = bool(options.get("refine_foreground", True))
if refine and src_h * src_w > _REFINE_PIXEL_LIMIT:
refine = False
self.last_warning = (
f"图像像素 {src_h * src_w / 1e6:.1f}MP 超过前景精修上限,已自动跳过该步骤"
)
if refine:
report("前景精修", 0.85)
rgb_out = np.clip(
np.rint(
imageops.refine_foreground(
rgb.astype(np.float32) / 255.0,
alpha_np,
int(options.get("blur_size") or 90),
int(options.get("blur_size_two") or 6),
)
* 255.0
),
0,
255,
).astype(np.uint8)
if background == "color":
cutout = imageops.compose_on_color(
rgb_out, alpha_np, imageops.parse_color(options.get("bg_color"))
)
else:
cutout = imageops.to_pil_rgba(rgb_out, alpha_np)
mask_pil = imageops.to_pil_mask(alpha_np) if want_mask else None
# 最终尺寸:按原图比例把最长边缩放到指定像素(0 = 保持原尺寸)
final_side = int(options.get("final_longest_side") or 0)
new_size = imageops.fit_longest_side(cutout.size, final_side)
if new_size != cutout.size:
cutout = cutout.resize(new_size, Image.LANCZOS)
if mask_pil is not None:
mask_pil = mask_pil.resize(new_size, Image.LANCZOS)
report("保存结果", 0.95)
# 前景占比:便于前端提示「模型可能没抠到东西」
coverage = float((alpha_np > 0.5).mean()) if alpha_np.size else 0.0
return {
"cutout": cutout,
"mask": mask_pil,
"width": cutout.width,
"height": cutout.height,
"source_width": src_w,
"source_height": src_h,
"input_size": [in_w, in_h],
"coverage": round(coverage, 4),
"model": info.name,
"model_path": info.path,
"backbone": info.backbone,
"arch": version,
"device": str(model_device),
"dtype": str(model_dtype),
"elapsed": round(time.perf_counter() - t0, 2),
"warning": self.last_warning,
"notes": self._notes(background, refine),
}
@staticmethod
def _notes(background: str, refine: bool) -> List[str]:
notes = []
if refine:
notes.append("已启用前景精修(去除边缘残留背景色)")
if background == "color":
notes.append("已合成到自定义背景色")
return notes
@staticmethod
def _forward(model, x, device: torch.device, model_dtype: torch.dtype, autocast_dtype):
"""前向一次,返回 sigmoid 后的低分辨率遮罩 (1,1,h,w)。"""
use_autocast = autocast_dtype is not None and str(device).startswith("cuda")
ctx = (
torch.autocast(device_type="cuda", dtype=autocast_dtype, enabled=True)
if use_autocast
else torch.autocast(device_type="cpu", enabled=False)
)
with torch.inference_mode(), ctx:
out = model(x)
pred = out[-1] if isinstance(out, (list, tuple)) else out
alpha = pred.sigmoid().float()
return alpha
# --------------------------------------------------------------------------- #
# 环境信息
# --------------------------------------------------------------------------- #
def describe_environment(engine: InferenceEngine) -> dict:
"""给前端展示的设备/运行环境信息。"""
try:
import torch as _torch # noqa: F401
import torchvision as _tv
tv_version = getattr(_tv, "__version__", "?")
except Exception: # pragma: no cover
tv_version = "?"
devices = []
if torch.cuda.is_available():
for i in range(torch.cuda.device_count()):
try:
props = torch.cuda.get_device_properties(i)
free, total = torch.cuda.mem_get_info(i)
devices.append(
{
"index": i,
"name": props.name,
"total_mem": round(total / 1024**3, 1),
"free_mem": round(free / 1024**3, 1),
"cc": f"{props.major}.{props.minor}",
}
)
except Exception:
devices.append({"index": i, "name": "CUDA 设备", "total_mem": 0, "free_mem": 0})
return {
"python": os.sys.version.split()[0],
"torch": torch.__version__,
"torchvision": tv_version,
"cuda": torch.cuda.is_available(),
"cuda_version": getattr(torch.version, "cuda", None),
"devices": devices,
"default_device": "cuda" if torch.cuda.is_available() else "cpu",
}
+67
View File
@@ -0,0 +1,67 @@
# -*- coding: utf-8 -*-
"""与推理无关的通用小工具:图片后缀集合、文件名清洗与尺寸换算。
刻意不 import torch / numpy —— 让路径校验、命名、尺寸计算这类纯逻辑可以脱离
推理环境被单独测试(见 :mod:`birefnet_web.batch`)。
"""
from __future__ import annotations
import os
from typing import Set, Tuple
#: 允许上传 / 扫描的图片后缀
IMAGE_SUFFIXES: Set[str] = {".png", ".jpg", ".jpeg", ".webp", ".bmp", ".tif", ".tiff", ".gif"}
#: 能装下 alpha 通道的后缀(透明输出只能用这些;jpg/bmp 装不下)
ALPHA_SUFFIXES: Set[str] = {".png", ".webp", ".tif", ".tiff"}
#: 文件名中不允许出现的字符(Windows 硬限制 + 控制字符)
BAD_NAME_CHARS = '<>:"/\\|?*\x00-\x1f'
def safe_stem(name: str, fallback: str = "image", maxlen: int = 80) -> str:
"""把外部来源的文件名清洗成安全的文件名主干(防路径穿越 + 去掉非法字符)。
Args:
name: 原始文件名或路径。
fallback: 清洗后为空时的兜底名字。
maxlen: 主干最大长度,超出截断。
Returns:
仅含安全字符的文件名主干(不含扩展名)。
"""
stem = os.path.splitext(os.path.basename(name or ""))[0]
stem = "".join("_" if ch in BAD_NAME_CHARS else ch for ch in stem).strip(" .")
stem = stem[:maxlen] or fallback
return stem
def fit_longest_side(size: Tuple[int, int], longest: int) -> Tuple[int, int]:
"""把 ``(宽, 高)`` 按原比例缩放到「最长边 == longest」。
与推理无关的纯几何换算,因此放在这里(而不是 imageops):结果尺寸只由
原尺寸和用户输入决定,可以脱离 torch 单测。
与 ``imageops.build_input_size(mode="longest")`` 的区别:那个算的是**网络
输入**尺寸,会对齐到 32 的倍数;这个算的是**最终成品**尺寸,不做对齐,
最长边精确等于用户填的值。
例:``(2000, 3000)`` + ``1440`` → ``(960, 1440)``;
``(3000, 2000)`` + ``1440`` → ``(1440, 960)``。
Args:
size: 原尺寸 ``(宽, 高)``,与 ``PIL.Image.size`` 一致。
longest: 目标最长边像素;``<= 0`` 表示保持原尺寸(用户默认值)。
Returns:
新尺寸 ``(宽, 高)``;无需缩放时原样返回。比例很小/很大时也至少 1 像素。
"""
w, h = int(size[0]), int(size[1])
if longest <= 0 or w <= 0 or h <= 0:
return w, h
current = max(w, h)
if current == longest:
return w, h
ratio = float(longest) / float(current)
return max(1, int(round(w * ratio))), max(1, int(round(h * ratio)))
+254
View File
@@ -0,0 +1,254 @@
# -*- coding: utf-8 -*-
"""图像预处理 / 后处理。
所有数值行为对齐 ComfyUI 节点 ``comfyui_birefnet_ll``:
* 预处理 —— 与 birefnetNode.ImagePreprocessor 相同的 Resize + ImageNet Normalize
* 遮罩还原 —— 等价 comfy.utils.common_upscale(F.interpolate,无抗锯齿)
* 前景精修 —— fast-foreground-estimation(Photoroom)的 box-blur 版本
"""
from __future__ import annotations
import os
from functools import lru_cache
from typing import Dict, Iterable, Optional, Sequence, Tuple
import numpy as np
import torch
from PIL import Image, ImageOps
# 与推理无关的通用工具(后缀集合、文件名清洗、尺寸换算)放在 fsutil 里,
# 这里重新导出,保持 `imageops.safe_stem` 这类既有调用点不变。
from .fsutil import ( # noqa: F401
ALPHA_SUFFIXES,
IMAGE_SUFFIXES,
fit_longest_side,
safe_stem,
)
try: # cv2 的 box blur 速度更快、背景更纯(与节点一致)
import cv2
_HAS_CV2 = True
except Exception: # pragma: no cover
cv2 = None # type: ignore[assignment]
_HAS_CV2 = False
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]
#: 与节点 interpolation_modes_mapping 对齐(torchvision InterpolationMode 的数值)
INTERPOLATIONS: Dict[str, int] = {
"nearest": 0,
"bilinear": 2,
"bicubic": 3,
"nearest-exact": 0,
}
#: 输出遮罩还原时允许的 F.interpolate 模式
UPSCALE_METHODS = ("bilinear", "nearest", "nearest-exact", "bicubic")
#: 预处理尺寸必须是 32 的倍数(Swin 骨干下采样 32 倍)
SIZE_ALIGN = 32
# --------------------------------------------------------------------------- #
# 基础转换
# --------------------------------------------------------------------------- #
def pil_to_rgb_array(image: Image.Image) -> np.ndarray:
"""PIL -> uint8 RGB 数组,同时按 EXIF 方向自动旋转。"""
if image.mode == "RGBA":
# 有透明通道时先合成到白底,避免 alpha 变黑
bg = Image.new("RGBA", image.size, (255, 255, 255, 255))
image = Image.alpha_composite(bg, image)
elif image.mode not in ("RGB", "L"):
image = image.convert("RGB")
image = ImageOps.exif_transpose(image)
if image.mode != "RGB":
image = image.convert("RGB")
return np.asarray(image, dtype=np.uint8)
def to_pil_rgba(rgb: np.ndarray, alpha: Optional[np.ndarray] = None) -> Image.Image:
"""uint8 RGB (+ float alpha[0,1]) -> PIL 图片。"""
if alpha is None:
return Image.fromarray(np.ascontiguousarray(rgb), "RGB")
a8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8)
rgba = np.dstack([rgb, a8])
return Image.fromarray(np.ascontiguousarray(rgba), "RGBA")
def to_pil_mask(alpha: np.ndarray) -> Image.Image:
"""float mask[0,1] -> 8bit 灰度图。"""
m8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8)
return Image.fromarray(np.ascontiguousarray(m8), "L")
def parse_color(value: object, default: Tuple[int, int, int] = (255, 255, 255)) -> Tuple[int, int, int]:
"""解析 ``#rgb`` / ``#rrggbb`` / ``(r,g,b)`` / int 为 RGB 三元组。"""
if value is None:
return default
if isinstance(value, (list, tuple)) and len(value) >= 3:
return tuple(int(np.clip(v, 0, 255)) for v in value[:3]) # type: ignore[return-value]
if isinstance(value, int):
return ((value >> 16) & 0xFF, (value >> 8) & 0xFF, value & 0xFF)
if isinstance(value, str):
s = value.strip().lstrip("#")
if len(s) == 3:
s = "".join(ch * 2 for ch in s)
if len(s) == 6:
try:
v = int(s, 16)
return ((v >> 16) & 0xFF, (v >> 8) & 0xFF, v & 0xFF)
except ValueError:
pass
return default
# --------------------------------------------------------------------------- #
# 输入尺寸
# --------------------------------------------------------------------------- #
def _align(v: float, align: int = SIZE_ALIGN, minimum: int = SIZE_ALIGN) -> int:
v = int(round(v / align) * align)
return max(minimum, v)
def build_input_size(
src_h: int,
src_w: int,
mode: str = "square",
width: int = 1024,
height: int = 1024,
longest_side: int = 1024,
) -> Tuple[int, int]:
"""计算网络输入尺寸 (h, w)。
mode:
square 固定 1024x1024(与节点默认一致,也是官方推荐分辨率)
longest 保持宽高比,长边 = longest_side,边长对齐到 32 的倍数
custom 使用自定义 width / height
"""
if src_h <= 0 or src_w <= 0:
return max(SIZE_ALIGN, int(height)), max(SIZE_ALIGN, int(width))
if mode == "custom":
h, w = _align(int(height)), _align(int(width))
elif mode == "longest":
target = max(SIZE_ALIGN, int(longest_side))
scale = target / float(max(src_h, src_w))
h, w = _align(src_h * scale), _align(src_w * scale)
else: # square
h, w = _align(int(height)), _align(int(width))
return h, w
# --------------------------------------------------------------------------- #
# 预处理
# --------------------------------------------------------------------------- #
@lru_cache(maxsize=32)
def _transform(size_hw: Tuple[int, int], method: str):
"""缓存 torchvision 变换(与节点 ImagePreprocessor 完全一致)。"""
from torchvision import transforms
interp = INTERPOLATIONS.get(method, 2)
return transforms.Compose(
[
transforms.Resize(size_hw, interpolation=interp),
transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
]
)
def preprocess(rgb: np.ndarray, size_hw: Tuple[int, int], method: str = "bilinear") -> torch.Tensor:
"""uint8 RGB (h,w,3) -> 归一化 float32 张量 (1,3,H,W)。"""
arr = np.asarray(rgb, dtype=np.float32) / 255.0 # 新建可写数组,避免 from_numpy 告警
tensor = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0)
return _transform(tuple(size_hw), method)(tensor)
def upscale_mask(mask_bchw: torch.Tensor, height: int, width: int, method: str = "bilinear") -> torch.Tensor:
"""把网络输出的低分辨率遮罩还原到原图尺寸(等价 comfy.utils.common_upscale)。"""
mode = method if method in UPSCALE_METHODS else "bilinear"
if tuple(mask_bchw.shape[-2:]) == (height, width):
return mask_bchw
return torch.nn.functional.interpolate(mask_bchw, size=(height, width), mode=mode)
def filter_mask(mask: torch.Tensor, threshold: float) -> torch.Tensor:
"""低于阈值的概率直接置零(与节点 util.filter_mask 一致)。"""
if threshold <= 0:
return mask
return mask * (mask > threshold).to(mask.dtype)
# --------------------------------------------------------------------------- #
# 前景精修(fast-foreground-estimation)
# --------------------------------------------------------------------------- #
def _box_blur(arr: np.ndarray, r: int) -> np.ndarray:
"""box blur;r 为核尺寸。cv2 可用时优先使用,保证与节点结果一致。
注意:cv2 对 (h, w, 1) 的输入会返回 (h, w),这里统一补齐通道维,
否则后续与 (h, w, 3) 广播会报错(原节点也是靠 `[:, :, None]` 兜住这一点)。
"""
r = int(max(1, r))
limit = 2 * min(arr.shape[0], arr.shape[1]) + 1
r = int(min(r, limit))
if _HAS_CV2:
out = cv2.blur(np.ascontiguousarray(arr, dtype=np.float32), (r, r))
return out[..., None] if (out.ndim == 2 and arr.ndim == 3) else out
# 兜底:torchvision 高斯模糊
from torchvision.transforms import functional as TF
if r % 2 == 0:
r += 1
t = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0).float()
out = TF.gaussian_blur(t, r) if r > 1 else t
return out.squeeze(0).permute(1, 2, 0).numpy()
def _fb_step(image: np.ndarray, F: np.ndarray, B: np.ndarray, alpha: np.ndarray, r: int):
a = alpha[:, :, None] if alpha.ndim == 2 else alpha
blurred_alpha = _box_blur(a, r)
blurred_F = _box_blur(F * a, r) / (blurred_alpha + 1e-5)
blurred_B = _box_blur(B * (1.0 - a), r) / ((1.0 - blurred_alpha) + 1e-5)
out = blurred_F + a * (image - a * blurred_F - (1.0 - a) * blurred_B)
return np.clip(out, 0.0, 1.0), blurred_B
def refine_foreground(rgb01: np.ndarray, alpha: np.ndarray, r1: int = 90, r2: int = 6) -> np.ndarray:
"""估计前景色,消除半透明边缘残留的原背景色。
参考 https://github.com/Photoroom/fast-foreground-estimation
"""
image = np.ascontiguousarray(rgb01, dtype=np.float32)
a = alpha.astype(np.float32)
F, blur_B = _fb_step(image, image, image, a, r1)
F2, _ = _fb_step(image, F, blur_B, a, r2)
return np.clip(F2, 0.0, 1.0)
# --------------------------------------------------------------------------- #
# 合成输出
# --------------------------------------------------------------------------- #
def compose_on_color(rgb: np.ndarray, alpha: np.ndarray, color: Sequence[int]) -> Image.Image:
"""按遮罩把前景合成到纯色背景上,返回 RGB 图。"""
a = alpha[:, :, None].astype(np.float32)
bg = np.array(color, dtype=np.float32).reshape(1, 1, 3)
out = rgb.astype(np.float32) * a + bg * (1.0 - a)
return Image.fromarray(np.clip(np.rint(out), 0, 255).astype(np.uint8), "RGB")
def save_image(image: Image.Image, path: str, quality: int = 95) -> str:
"""保存图片(自动建目录)。PNG/WebP 支持透明,JPEG 不支持。"""
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
ext = os.path.splitext(path)[1].lower()
if ext in (".jpg", ".jpeg"):
if image.mode == "RGBA":
bg = Image.new("RGBA", image.size, (255, 255, 255, 255))
image = Image.alpha_composite(bg, image).convert("RGB")
image.save(path, quality=quality, subsampling=0, optimize=True)
elif ext == ".webp":
image.save(path, quality=quality, method=4)
else:
image.save(path, optimize=False)
return path
+937
View File
@@ -0,0 +1,937 @@
# -*- coding: utf-8 -*-
"""HTTP 服务:静态页面 + REST API + 串行任务队列。
刻意不依赖任何 Web 框架(Gradio / FastAPI / Flask 都不需要):
* 传输层 —— 标准库 http.server.ThreadingHTTPServer
* 表单解析 —— 自研 multipart/form-data 解析(`cgi` 模块在 Python 3.13 已移除)
* 并发模型 —— 单 worker 线程串行消费任务,避免 GPU 显存被打爆
"""
from __future__ import annotations
import io
import json
import os
import queue
import threading
import time
import traceback
import uuid
import zipfile
from dataclasses import dataclass, field
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Sequence, Tuple
from urllib.parse import parse_qs, unquote, urlparse
from PIL import Image
from . import APP_NAME, __version__, batch, imageops
from .engine import InferenceEngine, ModelRegistry, describe_environment
# --------------------------------------------------------------------------- #
# 常量
# --------------------------------------------------------------------------- #
MAX_UPLOAD_BYTES = 512 * 1024 * 1024 # 单次请求体上限
MIME_TYPES = {
".html": "text/html; charset=utf-8",
".css": "text/css; charset=utf-8",
".js": "application/javascript; charset=utf-8",
".json": "application/json; charset=utf-8",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".webp": "image/webp",
".bmp": "image/bmp",
".tif": "image/tiff",
".tiff": "image/tiff",
".gif": "image/gif",
".svg": "image/svg+xml",
".ico": "image/x-icon",
".woff2": "font/woff2",
}
#: 允许上传 / 扫描的图片后缀(与批量模块共用同一份定义,避免两处不一致)
IMAGE_SUFFIXES = imageops.IMAGE_SUFFIXES
DEFAULT_OPTIONS: Dict[str, object] = {
"model": "",
"device": "auto",
"dtype": "auto",
"arch": "auto",
"resolution_mode": "square",
"width": 1024,
"height": 1024,
"longest_side": 1024,
"upscale_method": "bilinear",
"mask_threshold": 0.0,
"refine_foreground": True,
"blur_size": 90,
"blur_size_two": 6,
"background": "transparent",
"bg_color": "#ffffff",
"output_mask": False,
"final_longest_side": 0,
}
class TaskCanceled(BaseException):
"""取消信号。
故意继承 BaseException:这样它会穿过引擎里 `except Exception` 的进度回调保护,
直达任务循环,实现阶段级取消。
"""
# --------------------------------------------------------------------------- #
# multipart/form-data 解析
# --------------------------------------------------------------------------- #
@dataclass
class FormPart:
name: str
filename: Optional[str]
content_type: str
data: bytes
def _parse_disposition(value: str) -> Tuple[str, Optional[str]]:
name, filename = "", None
for chunk in value.split(";"):
chunk = chunk.strip()
low = chunk.lower()
if low.startswith("name=") and name == "":
name = chunk[5:].strip().strip('"')
elif low.startswith("filename="):
raw = chunk[9:].strip().strip('"')
if raw:
filename = raw
return name, filename
def parse_multipart(body: bytes, boundary: bytes) -> List[FormPart]:
"""把一个 multipart/form-data 请求体拆成若干 FormPart。"""
parts: List[FormPart] = []
delim = b"--" + boundary
for raw in body.split(delim):
if not raw or raw in (b"--", b"--\r\n", b"\r\n"):
continue
if raw.startswith(b"--"): # 结束标记
break
if raw.startswith(b"\r\n"):
raw = raw[2:]
elif raw.startswith(b"\n"):
raw = raw[1:]
head, sep, data = raw.partition(b"\r\n\r\n")
if not sep:
head, sep, data = raw.partition(b"\n\n")
if not sep:
continue
if data.endswith(b"\r\n"):
data = data[:-2]
elif data.endswith(b"\n"):
data = data[:-1]
headers: Dict[str, str] = {}
for line in head.split(b"\r\n"):
if b":" not in line:
continue
key, _, val = line.partition(b":")
headers[key.strip().lower().decode("latin-1")] = (
val.strip().decode("utf-8", "replace")
)
name, filename = _parse_disposition(headers.get("content-disposition", ""))
if not name:
continue
parts.append(
FormPart(
name=name,
filename=filename,
content_type=headers.get("content-type", "application/octet-stream"),
data=data,
)
)
return parts
# --------------------------------------------------------------------------- #
# 任务模型
# --------------------------------------------------------------------------- #
@dataclass
class ImageJob:
index: int
filename: str
label: str
input_path: str
state: str = "pending" # pending | running | done | failed | canceled
stage: str = "等待中"
progress: float = 0.0
error: Optional[str] = None
width: int = 0
height: int = 0
source_width: int = 0
source_height: int = 0
coverage: float = 0.0
elapsed: float = 0.0
input_size: Tuple[int, int] = (0, 0)
warning: Optional[str] = None
outputs: Dict[str, str] = field(default_factory=dict)
#: 批量模式:写进用户输出目录的文件名(RMBG_<原主干>_<时间戳>.<后缀>)
output_name: Optional[str] = None
def to_dict(self, task_id: str) -> dict:
urls = {
kind: f"/api/tasks/{task_id}/file/{self.index}/{kind}" for kind in self.outputs
}
return {
"index": self.index,
"name": self.filename,
"label": self.label,
"state": self.state,
"stage": self.stage,
"progress": round(self.progress, 3),
"error": self.error,
"width": self.width,
"height": self.height,
"source_width": self.source_width,
"source_height": self.source_height,
"input_size": list(self.input_size),
"coverage": self.coverage,
"elapsed": self.elapsed,
"warning": self.warning,
"output_name": self.output_name,
"urls": urls,
}
@dataclass
class Task:
id: str
created: float
options: dict
images: List[ImageJob]
output_dir: str
#: upload = 网页上传单图(结果落在 outputs/<task_id>);batch = 目录批量(结果落在用户输出目录)
mode: str = "upload"
input_dir: Optional[str] = None
recursive: bool = False
state: str = "queued" # queued | running | done | partial | failed | canceled
stage: str = "排队中"
progress: float = 0.0
error: Optional[str] = None
finished_at: Optional[float] = None
cancel_requested: bool = False
def to_dict(self, include_options: bool = False, tail: int = 0) -> dict:
"""序列化任务。
Args:
include_options: 是否带上参数快照。
tail: 大于 0 时只回传最后 N 张图片的状态。
批量任务动辄几百张,前端轮询必须靠它把载荷压住
(处理是按序串行的,所以窗口外的一定已经是终态,不会漏更新)。
"""
total = len(self.images)
start = max(0, total - int(tail)) if tail and tail > 0 else 0
payload = {
"id": self.id,
"created": self.created,
"mode": self.mode,
"input_dir": self.input_dir,
"output_dir": self.output_dir,
"recursive": self.recursive,
"state": self.state,
"stage": self.stage,
"progress": round(self.progress, 3),
"error": self.error,
"elapsed": round((self.finished_at or time.time()) - self.created, 2),
"counts": self._counts(),
"images_total": total,
"images_from": start,
"images": [img.to_dict(self.id) for img in self.images[start:]],
}
if include_options:
payload["options"] = self.options
return payload
def _counts(self) -> dict:
out = {"total": len(self.images), "done": 0, "failed": 0, "pending": 0}
for img in self.images:
if img.state == "done":
out["done"] += 1
elif img.state == "failed":
out["failed"] += 1
elif img.state != "canceled":
out["pending"] += 1
return out
def refresh_progress(self) -> None:
if not self.images:
self.progress = 0.0
else:
self.progress = sum(i.progress for i in self.images) / len(self.images)
class TaskManager:
"""单 worker 串行任务队列。"""
def __init__(
self,
engine: InferenceEngine,
output_root: Path,
max_tasks: int = 50,
logger=None,
) -> None:
self.engine = engine
self.output_root = Path(output_root)
self.output_root.mkdir(parents=True, exist_ok=True)
self.max_tasks = max(1, int(max_tasks))
self.log = logger or (lambda msg: None)
self._tasks: Dict[str, Task] = {}
self._lock = threading.RLock()
self._queue: "queue.Queue[Optional[Task]]" = queue.Queue()
self._worker = threading.Thread(target=self._work_loop, name="birefnet-worker", daemon=True)
self._worker.start()
# -------------------------- 对外接口 -------------------------- #
def submit(self, uploads: Sequence[Tuple[str, bytes]], options: dict) -> Task:
if not uploads:
raise ValueError("没有收到任何图片")
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
output_dir = self.output_root / task_id
input_dir = output_dir / "input"
input_dir.mkdir(parents=True, exist_ok=True)
jobs: List[ImageJob] = []
for idx, (filename, data) in enumerate(uploads):
label = imageops.safe_stem(filename, fallback=f"image_{idx + 1}")
suffix = Path(filename or "").suffix.lower()
if suffix not in IMAGE_SUFFIXES:
suffix = ".png"
stored = input_dir / f"{idx + 1:03d}_{label}{suffix}"
stored.write_bytes(data)
jobs.append(ImageJob(index=idx, filename=filename or stored.name, label=label, input_path=str(stored)))
task = Task(id=task_id, created=time.time(), options=options, images=jobs, output_dir=str(output_dir))
with self._lock:
self._tasks[task_id] = task
self._prune_locked()
self._queue.put(task)
self.log(f"[任务 {task_id}] 已提交,共 {len(jobs)} 张图片")
return task
def submit_batch(
self,
input_dir: Path,
output_dir: Path,
options: dict,
*,
recursive: bool = False,
) -> Task:
"""按目录批量建任务:原图不拷贝,结果直接写进用户指定的输出目录。
Args:
input_dir: 已校验存在的图片目录。
output_dir: 已校验存在的输出目录(本方法不创建任何目录)。
options: 与单图模式完全相同的参数集合。
recursive: 是否包含子目录。
Returns:
已入队的新任务。
Raises:
ValueError: 目录下没有可处理的图片。
"""
# 输入 == 输出时不能把输出目录当跳过项,否则一张都扫不到
skip = [] if Path(input_dir) == Path(output_dir) else [output_dir]
files = batch.scan_images(input_dir, recursive=recursive, skip_dirs=skip)
if not files:
raise ValueError(f"目录下没有找到可处理的图片:{input_dir}")
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
jobs = [
ImageJob(
index=idx,
filename=path.name,
label=imageops.safe_stem(path.name, fallback=f"image_{idx + 1}"),
input_path=str(path),
)
for idx, path in enumerate(files)
]
task = Task(
id=task_id,
created=time.time(),
options=options,
images=jobs,
output_dir=str(output_dir),
mode="batch",
input_dir=str(input_dir),
recursive=recursive,
)
with self._lock:
self._tasks[task_id] = task
self._prune_locked()
self._queue.put(task)
self.log(f"[批量 {task_id}] 已提交,{len(jobs)} 张图片 → 输出目录 {output_dir}")
return task
def get(self, task_id: str) -> Optional[Task]:
with self._lock:
return self._tasks.get(task_id)
def list_tasks(self, limit: int = 20) -> List[dict]:
"""任务列表(含各自使用的参数,供前端恢复结果历史时直接显示参数)。"""
with self._lock:
tasks = sorted(self._tasks.values(), key=lambda t: t.created, reverse=True)[:limit]
return [t.to_dict(include_options=True) for t in tasks]
def cancel(self, task_id: str) -> bool:
task = self.get(task_id)
if task is None or task.finished_at is not None:
return False
task.cancel_requested = True
if task.state == "queued":
task.state = "canceled"
task.stage = "已取消"
task.finished_at = time.time()
self.log(f"[任务 {task_id}] 收到取消请求")
return True
def _prune_locked(self) -> None:
if len(self._tasks) <= self.max_tasks:
return
ordered = sorted(self._tasks.values(), key=lambda t: t.created)
for stale in ordered[: len(self._tasks) - self.max_tasks]:
if stale.finished_at is not None:
self._tasks.pop(stale.id, None)
# -------------------------- worker -------------------------- #
def _work_loop(self) -> None:
while True:
task = self._queue.get()
if task is None:
return
if task.cancel_requested:
continue
try:
self._process(task)
except TaskCanceled: # 正常取消,不需要堆栈
task.state = "canceled"
task.stage = "已取消"
task.finished_at = time.time()
except BaseException as exc: # noqa: BLE001 - worker 必须兜住一切
self.log(f"[任务 {task.id}] 异常终止:{exc}")
traceback.print_exc()
task.state = "failed"
task.error = str(exc)
task.stage = "失败"
task.finished_at = time.time()
def _process(self, task: Task) -> None:
task.state = "running"
task.stage = "开始处理"
self.log(f"[任务 {task.id}] 开始处理,参数:{json.dumps(task.options, ensure_ascii=False)}")
for job in task.images:
if task.cancel_requested:
if job.state in ("pending", "running"):
job.state = "canceled"
job.stage = "已取消"
job.progress = 1.0
continue
try:
self._process_one(task, job)
except TaskCanceled:
# 当前图片已被标记为 canceled;剩余图片由上面的分支收尾
continue
counts = task._counts()
task.refresh_progress()
task.progress = 1.0
task.finished_at = time.time()
if task.cancel_requested:
task.state = "canceled"
task.stage = "已取消"
elif counts["failed"] and counts["done"]:
task.state = "partial"
task.stage = "部分完成"
elif counts["failed"]:
task.state = "failed"
task.stage = "失败"
else:
task.state = "done"
task.stage = "完成"
self.log(
f"[任务 {task.id}] 结束:{task.state},成功 {counts['done']} / 失败 {counts['failed']},"
f"耗时 {task.finished_at - task.created:.1f}s"
)
def _process_one(self, task: Task, job: ImageJob) -> None:
job.state = "running"
job.progress = 0.02
task.refresh_progress()
def on_progress(stage: str, frac: float) -> None:
if task.cancel_requested:
raise TaskCanceled()
job.stage = stage
job.progress = 0.05 + 0.9 * float(frac)
task.stage = f"{job.label} · {stage}"
task.refresh_progress()
try:
with Image.open(job.input_path) as im:
im.load()
pil = im.copy()
job.source_width, job.source_height = pil.size
result = self.engine.remove_background(pil, task.options, on_progress)
job.outputs = self._store_outputs(task, job, result)
job.width = result["width"]
job.height = result["height"]
job.coverage = result["coverage"]
job.elapsed = result["elapsed"]
job.input_size = tuple(result["input_size"]) # type: ignore[assignment]
job.warning = result["warning"]
job.state = "done"
job.stage = "完成"
job.progress = 1.0
target = f" → {job.output_name}" if job.output_name else ""
self.log(
f"[{'批量' if task.mode == 'batch' else '任务'} {task.id}] {job.label}{target} 完成:"
f"{job.width}x{job.height},前景占比 {job.coverage:.1%},耗时 {job.elapsed}s"
)
except TaskCanceled:
job.state = "canceled"
job.stage = "已取消"
job.progress = 1.0
raise
except Exception as exc: # 单张失败不影响其余图片
job.state = "failed"
job.stage = "失败"
job.progress = 1.0
job.error = f"{type(exc).__name__}: {exc}"
self.log(f"[任务 {task.id}] {job.label} 处理失败:{job.error}")
traceback.print_exc()
finally:
task.refresh_progress()
def _store_outputs(self, task: Task, job: ImageJob, result: dict) -> Dict[str, str]:
"""把单张结果落盘,返回 ``kind -> 路径``。
* 单图模式:沿用 ``outputs/<task_id>/<序号>_<主干>_cutout.png``。
* 批量模式:写进用户指定的输出目录,文件名按需求固定为
``RMBG_<原主干(≤20)>_<Unix 时间戳>.<原后缀>``;同名时递增 ``_1``
避让而不是覆盖(见 :mod:`birefnet_web.batch`)。
"""
outputs: Dict[str, str] = {}
if task.mode == "batch":
out_dir = Path(task.output_dir)
name = batch.build_output_name(
Path(job.input_path),
int(time.time()),
str(task.options.get("background") or "transparent"),
)
target = batch.unique_path(out_dir, name)
imageops.save_image(result["cutout"], str(target))
job.output_name = target.name
outputs["cutout"] = str(target)
if result["mask"] is not None:
mask_target = batch.unique_path(out_dir, f"{target.stem}_mask.png")
imageops.save_image(result["mask"], str(mask_target))
outputs["mask"] = str(mask_target)
else:
cutout_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_cutout.png")
imageops.save_image(result["cutout"], cutout_png)
outputs["cutout"] = cutout_png
if result["mask"] is not None:
mask_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_mask.png")
imageops.save_image(result["mask"], mask_png)
outputs["mask"] = mask_png
outputs["original"] = job.input_path
return outputs
# --------------------------------------------------------------------------- #
# HTTP 处理器
# --------------------------------------------------------------------------- #
class WebUIHandler(BaseHTTPRequestHandler):
server_version = f"BiRefNetWebUI/{__version__}"
protocol_version = "HTTP/1.1"
# ------------------------- 基础工具 ------------------------- #
@property
def app(self) -> "WebUIServer": # type: ignore[override]
return self.server # type: ignore[return-value]
def log_message(self, fmt: str, *args) -> None: # noqa: A003
path = getattr(self, "path", "")
# 轮询与静态资源太吵,不打印
if path.startswith("/api/tasks/") or path.startswith("/static/") or path == "/favicon.ico":
return
super().log_message(fmt, *args)
def _send(self, status: int, body: bytes, content_type: str, extra: Optional[dict] = None) -> None:
try:
self.send_response(status)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
for k, v in (extra or {}).items():
self.send_header(k, v)
self.end_headers()
if self.command != "HEAD":
self.wfile.write(body)
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
def _send_json(self, payload: object, status: int = 200) -> None:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
self._send(status, body, "application/json; charset=utf-8")
def _error(self, status: int, message: str) -> None:
self._send_json({"ok": False, "error": message}, status=status)
def _read_body(self) -> bytes:
try:
length = int(self.headers.get("Content-Length") or 0)
except ValueError:
length = 0
if length <= 0:
return b""
if length > MAX_UPLOAD_BYTES:
raise ValueError(
f"请求体过大({length / 1024 / 1024:.1f} MB),上限 {MAX_UPLOAD_BYTES / 1024 / 1024:.0f} MB"
)
return self.rfile.read(length)
# ------------------------- 路由 ------------------------- #
def do_GET(self) -> None: # noqa: N802
try:
parsed = urlparse(self.path)
path = unquote(parsed.path)
query = parse_qs(parsed.query)
if path in ("/", "/index.html"):
return self._serve_static("index.html")
if path.startswith("/static/"):
return self._serve_static(path[len("/static/"):])
if path == "/favicon.ico":
return self._send(204, b"", "image/x-icon")
if path == "/api/health":
return self._send_json({"ok": True, "app": APP_NAME, "version": __version__})
if path == "/api/state":
return self._api_state()
if path == "/api/models":
return self._send_json({"ok": True, "models": [m.to_dict() for m in self.app.registry.list()]})
if path == "/api/tasks":
return self._send_json({"ok": True, "tasks": self.app.tasks.list_tasks()})
parts = path.strip("/").split("/")
# /api/tasks/<id>[/...]
if len(parts) >= 3 and parts[0] == "api" and parts[1] == "tasks":
task_id = parts[2]
if len(parts) == 3:
task = self.app.tasks.get(task_id)
if task is None:
return self._error(404, "任务不存在")
return self._send_json({"ok": True, "task": self._task_payload(task, query)})
if len(parts) == 6 and parts[3] == "file":
return self._serve_result(task_id, parts[4], parts[5])
if len(parts) == 4 and parts[3] == "zip":
return self._serve_zip(task_id, query)
return self._error(404, f"未知路径 {path}")
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
except Exception as exc: # pragma: no cover
traceback.print_exc()
self._error(500, f"{type(exc).__name__}: {exc}")
def do_HEAD(self) -> None: # noqa: N802
self.do_GET()
def do_POST(self) -> None: # noqa: N802
try:
parsed = urlparse(self.path)
path = unquote(parsed.path)
if path == "/api/models/reload":
models = self.app.registry.scan()
return self._send_json({"ok": True, "models": [m.to_dict() for m in models]})
if path == "/api/engine/unload":
self.app.engine.unload()
return self._send_json({"ok": True})
if path == "/api/shutdown":
self._send_json({"ok": True, "message": "服务正在关闭"})
threading.Thread(target=self.app.stop_soon, daemon=True).start()
return
if path == "/api/tasks":
return self._create_task()
if path == "/api/batch":
return self._create_batch_task()
parts = path.strip("/").split("/")
if len(parts) == 4 and parts[0] == "api" and parts[1] == "tasks" and parts[3] == "cancel":
ok = self.app.tasks.cancel(parts[2])
return self._send_json({"ok": ok})
return self._error(404, f"未知路径 {path}")
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
except Exception as exc: # pragma: no cover
traceback.print_exc()
self._error(500, f"{type(exc).__name__}: {exc}")
# ------------------------- 具体处理 ------------------------- #
@staticmethod
def _int_param(query: dict, key: str, default: int) -> int:
"""从 query(parse_qs 的结果)里安全地取一个整数参数。"""
try:
return int(str((query.get(key) or [default])[0]))
except (TypeError, ValueError):
return default
def _task_payload(self, task: Task, query: dict) -> dict:
"""任务详情;``?tail=N`` 只回传最后 N 张(批量任务轮询用,压住载荷)。"""
return task.to_dict(include_options=True, tail=max(0, self._int_param(query, "tail", 0)))
def _merge_options(self, incoming: object) -> dict:
"""把前端传来的参数合并到默认参数上(只接受白名单键)。"""
options = dict(DEFAULT_OPTIONS)
if isinstance(incoming, dict):
options.update({k: v for k, v in incoming.items() if k in DEFAULT_OPTIONS})
return options
def _pin_model(self, options: dict) -> Optional[str]:
"""确保 ``options['model']`` 指向一个真实存在的权重。
Returns:
出错时返回给人看的错误信息;成功返回 None。
"""
models = self.app.registry.list()
if not models:
return "模型目录中没有找到任何权重文件(*.safetensors / *.pth)"
options["model"] = options.get("model") or self.app.default_model_key(
[m.to_dict() for m in models]
)
try:
self.app.registry.get(str(options["model"]))
except KeyError:
options["model"] = models[0].key
return None
def _api_state(self) -> None:
app = self.app
models = [m.to_dict() for m in app.registry.list()]
default_model = app.default_model_key(models)
self._send_json(
{
"ok": True,
"app": APP_NAME,
"version": __version__,
"environment": describe_environment(app.engine),
"node_dir": str(app.node_dir),
"model_dirs": [str(d) for d in app.registry.model_dirs],
"output_dir": str(app.output_root),
"project_root": str(app.project_root),
"models": models,
"defaults": {**DEFAULT_OPTIONS, "model": default_model},
}
)
def _create_task(self) -> None:
content_type = self.headers.get("Content-Type", "")
if "multipart/form-data" not in content_type:
return self._error(400, "Content-Type 必须是 multipart/form-data")
boundary = None
for chunk in content_type.split(";"):
chunk = chunk.strip()
if chunk.lower().startswith("boundary="):
boundary = chunk[9:].strip().strip('"').encode("latin-1")
if not boundary:
return self._error(400, "缺少 multipart boundary")
try:
body = self._read_body()
except ValueError as exc:
return self._error(413, str(exc))
options = dict(DEFAULT_OPTIONS)
uploads: List[Tuple[str, bytes]] = []
for part in parse_multipart(body, boundary):
if part.name == "options":
try:
options = self._merge_options(json.loads(part.data.decode("utf-8") or "{}"))
except json.JSONDecodeError:
return self._error(400, "options 字段不是合法 JSON")
elif part.name in ("files", "images", "file"):
if part.data:
uploads.append((part.filename or f"upload_{len(uploads) + 1}.png", part.data))
if not uploads:
return self._error(400, "没有收到任何图片,请选择 PNG / JPG / WebP 文件")
error = self._pin_model(options)
if error:
return self._error(400, error)
try:
task = self.app.tasks.submit(uploads, options)
except ValueError as exc:
return self._error(400, str(exc))
self._send_json({"ok": True, "task_id": task.id, "count": len(uploads)})
def _create_batch_task(self) -> None:
"""POST /api/batch —— 目录批量抠图(JSON 传路径,不上传文件)。
请求体:
{
"input_dir": "F:\\\\photos\\\\待抠图", # 必填,图片所在目录
"output_dir": "F:\\\\photos\\\\results", # 可空 → 用项目 outputs/
"recursive": false, # 是否含子目录
"options": { …与单图模式相同的参数… }
}
两个目录都必须**已经存在**:路径不对就报错,本接口绝不会创建目录。
"""
try:
payload = json.loads(self._read_body().decode("utf-8") or "{}")
except (ValueError, UnicodeDecodeError):
return self._error(400, "请求体不是合法 JSON")
if not isinstance(payload, dict):
return self._error(400, "请求体必须是 JSON 对象")
base = self.app.project_root
raw_out = str(payload.get("output_dir") or "").strip() or str(self.app.output_root)
try:
input_dir = batch.check_input_dir(payload.get("input_dir"), base)
output_dir = batch.check_output_dir(raw_out, base)
except batch.BatchPathError as exc:
return self._error(400, str(exc))
options = self._merge_options(payload.get("options"))
error = self._pin_model(options)
if error:
return self._error(400, error)
try:
task = self.app.tasks.submit_batch(
input_dir, output_dir, options, recursive=bool(payload.get("recursive"))
)
except ValueError as exc:
return self._error(400, str(exc))
self._send_json(
{
"ok": True,
"task_id": task.id,
"mode": "batch",
"count": len(task.images),
"input_dir": str(input_dir),
"output_dir": str(output_dir),
}
)
def _task_file(self, task_id: str, index: int, kind: str) -> Optional[str]:
task = self.app.tasks.get(task_id)
if task is None:
return None
for job in task.images:
if job.index == index:
return job.outputs.get(kind)
return None
def _serve_result(self, task_id: str, idx_part: str, kind: str) -> None:
try:
index = int(idx_part)
except ValueError:
return self._error(400, "图片序号非法")
if kind not in ("cutout", "mask", "original"):
return self._error(400, f"未知的结果类型 {kind!r}")
path = self._task_file(task_id, index, kind)
if not path or not os.path.isfile(path):
return self._error(404, "结果文件不存在")
suffix = Path(path).suffix.lower()
with open(path, "rb") as fh:
data = fh.read()
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"))
def _serve_zip(self, task_id: str, query: dict) -> None:
task = self.app.tasks.get(task_id)
if task is None:
return self._error(404, "任务不存在")
kinds = (query.get("kind") or ["cutout"])[0].split(",")
want = [k for k in kinds if k in ("cutout", "mask", "original")] or ["cutout"]
buffer = io.BytesIO()
added = 0
with zipfile.ZipFile(buffer, "w", zipfile.ZIP_DEFLATED) as zf:
for job in task.images:
for kind in want:
path = job.outputs.get(kind)
if path and os.path.isfile(path):
zf.write(path, arcname=os.path.basename(path))
added += 1
if added == 0:
return self._error(404, "没有可打包的结果文件")
self._send(
200,
buffer.getvalue(),
"application/zip",
extra={"Content-Disposition": f'attachment; filename="birefnet_{task_id}.zip"'},
)
def _serve_static(self, rel_path: str) -> None:
web_dir = self.app.web_dir.resolve()
target = (web_dir / rel_path.replace("\\", "/")).resolve()
try:
target.relative_to(web_dir) # 防路径穿越
except ValueError:
return self._error(403, "非法路径")
if not target.is_file():
return self._error(404, f"资源不存在:{rel_path}")
suffix = target.suffix.lower()
with open(target, "rb") as fh:
data = fh.read()
cache = "no-cache" if suffix in (".html", ".js", ".css") else "public, max-age=3600"
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"), {"Cache-Control": cache})
# --------------------------------------------------------------------------- #
# 服务容器
# --------------------------------------------------------------------------- #
class WebUIServer(ThreadingHTTPServer):
daemon_threads = True
allow_reuse_address = True
def __init__(
self,
address: Tuple[str, int],
registry: ModelRegistry,
engine: InferenceEngine,
tasks: TaskManager,
web_dir: Path,
output_root: Path,
node_dir: Path,
logger=None,
project_root: Optional[Path] = None,
) -> None:
super().__init__(address, WebUIHandler)
self.registry = registry
self.engine = engine
self.tasks = tasks
self.web_dir = Path(web_dir)
self.output_root = Path(output_root)
self.node_dir = Path(node_dir)
#: 相对路径型用户输入(批量处理的目录)以此为准
self.project_root = Path(project_root) if project_root else self.web_dir.parent
self.log = logger or (lambda msg: None)
def default_model_key(self, models: Optional[Iterable[dict]] = None) -> str:
models = list(models if models is not None else (m.to_dict() for m in self.registry.list()))
if not models:
return ""
for preferred in ("Portrait", "General", "General-HR"):
for m in models:
if m["name"] == preferred:
return m["key"]
return models[0]["key"]
def stop_soon(self) -> None:
time.sleep(0.4)
self.log("收到关闭指令,服务即将退出")
self.shutdown()