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