Files
BiRefNet_WebUI/birefnet_web/engine.py
T
2026-10-08 09:46:47 +08:00

500 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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",
}