263 lines
9.5 KiB
Python
263 lines
9.5 KiB
Python
# -*- 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
|