# -*- 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", }