feat: initial project setup
This commit is contained in:
@@ -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"
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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)))
|
||||
@@ -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
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user