Files
2026-10-08 09:46:47 +08:00

358 lines
14 KiB
Python
Raw Permalink 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 -*-
"""BiRefNet WebUI 启动入口。
本项目**自包含**:模型代码、权重、Python 运行时都在项目目录内,不再需要 ComfyUI。
F:\\BiRefNet_WebUI\\
├─ Webui.py
├─ python\\ 自带 Python 3.12 运行时(torch + CUDA)
├─ models\\ BiRefNet 权重(*.safetensors)
├─ vendor\\comfyui_birefnet_ll\\ 模型代码
└─ birefnet_web\\ web\\ outputs\\
启动方式(任选其一):
run.bat 双击即可
python\\python.exe Webui.py 用自带运行时
python Webui.py 用系统 Python(缺依赖时会自动切到自带运行时)
常用参数:
--port 7861 监听端口
--host 0.0.0.0 允许局域网访问
--device cpu 强制 CPU 推理
--dtype float32 强制单精度(默认 GPU 上 fp32 权重 + fp16 autocast)
--model-dir D:\\models 追加模型目录(可重复)
--node-dir <path> 指定模型代码目录(默认用项目内 vendor/)
--no-browser 启动后不自动打开浏览器
--check 仅自检环境后退出
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import threading
import time
import webbrowser
from pathlib import Path
from typing import Dict, List, Optional
PROJECT_ROOT = Path(__file__).resolve().parent
CONFIG_PATH = PROJECT_ROOT / "config.json"
# 内置 Python 带有 python3xx._pth(隔离模式),脚本所在目录不会自动进入 sys.path,
# 这里显式补上,保证无论用哪个解释器都能 import 到本项目包。
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
#: 项目自带的独立 Python 运行时
BUNDLED_PYTHON = PROJECT_ROOT / "python" / "python.exe"
#: 防止自动切换解释器时无限递归
_RELAUNCH_FLAG = "BIREFNET_RELAUNCHED"
#: 运行必需的三方包 -> pip 包名
_REQUIRED = (
("torch", "torch"),
("torchvision", "torchvision"),
("numpy", "numpy"),
("PIL", "pillow"),
("safetensors", "safetensors"),
("timm", "timm"),
("einops", "einops"),
("kornia", "kornia"),
)
CONFIG_DEFAULTS: Dict[str, object] = {
"host": "127.0.0.1",
"port": 7861,
"node_dir": "",
"model_dirs": [],
"output_dir": "outputs",
"device": "auto",
"dtype": "auto",
"max_cached_models": 1,
"max_tasks": 50,
"open_browser": True,
}
# --------------------------------------------------------------------------- #
# 依赖自检 / 解释器自动切换
# --------------------------------------------------------------------------- #
def missing_packages() -> List[str]:
"""返回缺失的 pip 包名列表。"""
missing: List[str] = []
for mod, pkg in _REQUIRED:
try:
__import__(mod)
except Exception:
missing.append(pkg)
return missing
def bundled_python() -> Optional[Path]:
"""项目自带运行时是否可用。"""
return BUNDLED_PYTHON if BUNDLED_PYTHON.is_file() else None
def relaunch_with_bundled_python(missing: List[str]) -> None:
"""当前解释器缺依赖、而项目自带运行时可用时,直接把进程换成自带运行时。"""
if not missing or os.environ.get(_RELAUNCH_FLAG):
return
exe = bundled_python()
if exe is None:
return
try:
if Path(sys.executable).resolve() == exe.resolve():
return
except OSError:
pass
os.environ[_RELAUNCH_FLAG] = "1"
print(f"[环境] 当前解释器缺少 {', '.join(missing)},自动切换到项目自带运行时:{exe}", flush=True)
try:
os.execv(str(exe), [str(exe), str(Path(__file__).resolve()), *sys.argv[1:]])
except OSError as exc: # 切换失败就继续走下面的报错分支
print(f"[环境] 切换失败:{exc}", file=sys.stderr, flush=True)
def preflight() -> None:
"""尽早给出人话的依赖缺失提示,而不是一堆 ImportError 堆栈。"""
missing = missing_packages()
if missing:
relaunch_with_bundled_python(missing)
if missing:
print("缺少依赖:" + ", ".join(missing), file=sys.stderr)
if bundled_python() is None:
print(
f"项目自带运行时不存在({BUNDLED_PYTHON})。\n"
"请从 ComfyUI 便携版的 python_embeded 目录复制一份到项目的 python\\ 下,"
"或用你自己的解释器安装依赖:",
file=sys.stderr,
)
print(f' "{sys.executable}" -m pip install -r requirements.txt', file=sys.stderr)
sys.exit(2)
if not _has_cv2():
print("[提示] 未检测到 opencv-python,前景精修将退化为高斯模糊实现(结果略有差异)")
def _has_cv2() -> bool:
try:
import cv2 # noqa: F401
return True
except Exception:
return False
# --------------------------------------------------------------------------- #
# 配置
# --------------------------------------------------------------------------- #
def load_config(path: Path) -> Dict[str, object]:
config = dict(CONFIG_DEFAULTS)
if path.is_file():
try:
data = json.loads(path.read_text(encoding="utf-8"))
if isinstance(data, dict):
config.update(data)
except (json.JSONDecodeError, OSError) as exc:
print(f"[警告] 读取 {path} 失败,使用默认配置:{exc}")
return config
def save_config(path: Path, config: Dict[str, object]) -> None:
try:
path.write_text(json.dumps(config, ensure_ascii=False, indent=2), encoding="utf-8")
except OSError as exc: # pragma: no cover
print(f"[警告] 写入 {path} 失败:{exc}")
def parse_args(argv: Optional[List[str]] = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
prog="Webui.py",
description="BiRefNet WebUI —— 自包含的网页版智能抠图服务",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--host", default=None, help="监听地址,0.0.0.0 可局域网访问")
parser.add_argument("--port", type=int, default=None, help="监听端口")
parser.add_argument("--device", default=None, choices=["auto", "cpu", "cuda"], help="推理设备")
parser.add_argument("--dtype", default=None, choices=["auto", "float32", "float16", "bfloat16"], help="推理精度")
parser.add_argument("--model-dir", action="append", default=None, help="追加模型目录,可重复指定")
parser.add_argument("--node-dir", default=None, help="模型代码目录(默认用项目内 vendor/)")
parser.add_argument("--output-dir", default=None, help="结果输出目录")
parser.add_argument("--max-cached-models", type=int, default=None, help="显存中最多缓存的模型数量")
parser.add_argument("--max-tasks", type=int, default=None, help="内存中保留的任务记录数量")
parser.add_argument("--config", default=str(CONFIG_PATH), help="配置文件路径")
parser.add_argument("--no-browser", action="store_true", help="启动后不自动打开浏览器")
parser.add_argument("--save-config", action="store_true", help="把本次参数写回配置文件")
parser.add_argument("--check", action="store_true", help="仅做环境自检并打印可用模型,然后退出")
return parser.parse_args(argv)
def merge_config(args: argparse.Namespace) -> Dict[str, object]:
config_path = Path(args.config).expanduser().resolve()
config = load_config(config_path)
if args.host is not None:
config["host"] = args.host
if args.port is not None:
config["port"] = args.port
if args.device is not None:
config["device"] = args.device
if args.dtype is not None:
config["dtype"] = args.dtype
if args.output_dir is not None:
config["output_dir"] = args.output_dir
if args.max_cached_models is not None:
config["max_cached_models"] = args.max_cached_models
if args.max_tasks is not None:
config["max_tasks"] = args.max_tasks
if args.model_dir:
config["model_dirs"] = list(dict.fromkeys(list(config.get("model_dirs") or []) + args.model_dir))
if args.no_browser:
config["open_browser"] = False
# 命令行显式指定的代码目录优先级最高,单独存一个键,避免被 config.json 里的旧值盖住
config["_cli_node_dir"] = args.node_dir
config["_path"] = str(config_path)
config["_first_run"] = not config_path.is_file()
return config
# --------------------------------------------------------------------------- #
# 主流程
# --------------------------------------------------------------------------- #
def build_app(config: Dict[str, object]):
from birefnet_web import compat
from birefnet_web.engine import InferenceEngine, ModelRegistry
from birefnet_web.server import TaskManager, WebUIServer
# 1) 模型代码目录:命令行 > 环境变量 > 项目内 vendor/ > 旧配置 > 外部 ComfyUI
node_dir = compat.find_node_dir(config.get("_cli_node_dir"), config.get("node_dir") or None)
# 2) 权重目录:项目内 models/ 优先
model_dirs: List[Path] = []
for d in compat.default_model_dirs(node_dir):
if d not in model_dirs:
model_dirs.append(d)
for item in config.get("model_dirs") or []:
if item:
p = Path(str(item)).expanduser()
if p not in model_dirs:
model_dirs.append(p)
registry = ModelRegistry(model_dirs)
models = registry.scan()
output_root = Path(str(config.get("output_dir") or "outputs")).expanduser()
if not output_root.is_absolute():
output_root = PROJECT_ROOT / output_root
engine = InferenceEngine(registry, node_dir=node_dir, max_cached=int(config.get("max_cached_models") or 1))
logger = lambda msg: print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True) # noqa: E731
tasks = TaskManager(engine, output_root, max_tasks=int(config.get("max_tasks") or 50), logger=logger)
server = WebUIServer(
(str(config.get("host") or "127.0.0.1"), int(config.get("port") or 7861)),
registry=registry,
engine=engine,
tasks=tasks,
web_dir=PROJECT_ROOT / "web",
output_root=output_root,
node_dir=node_dir,
logger=logger,
project_root=PROJECT_ROOT,
)
return server, registry, models, node_dir, model_dirs, output_root
def _rel_or_abs(path: Path) -> str:
"""项目内的路径显示成相对路径,读起来更舒服。"""
try:
return str(path.resolve().relative_to(PROJECT_ROOT))
except (ValueError, OSError):
return str(path)
def print_banner(server, models, node_dir, model_dirs, output_root, config) -> None:
host, port = server.server_address[0], server.server_address[1]
shown = "127.0.0.1" if host in ("0.0.0.0", "::") else host
url = f"http://{shown}:{port}/"
from birefnet_web.engine import describe_environment
env = describe_environment(server.engine)
line = "─" * 66
print(line)
print(" BiRefNet WebUI · 智能抠图")
print(line)
print(f" 访问地址 {url}")
if host in ("0.0.0.0", "::"):
print(f" 局域网 同一网络下用 http://<本机IP>:{port}/ 访问")
print(f" 运行时 {_rel_or_abs(Path(sys.executable))} (python {env['python']})")
print(f" 推理库 torch {env['torch']} / torchvision {env['torchvision']}")
if env["cuda"]:
for dev in env["devices"]:
print(f" GPU [{dev['index']}] {dev['name']} {dev['free_mem']}G 空闲 / {dev['total_mem']}G")
else:
print(" GPU 未检测到 CUDA,将使用 CPU(较慢)")
print(f" 模型代码 {_rel_or_abs(node_dir)}")
for d in model_dirs:
mark = "✔" if Path(d).is_dir() else "✘ 不存在"
print(f" 模型目录 {_rel_or_abs(Path(d))} {mark}")
print(f" 可用模型 {len(models)} 个" + (f":{', '.join(m.name for m in models)}" if models else "(未找到权重文件)"))
if not models:
print(" ↳ 请把 *.safetensors 放到上面的模型目录")
print(f" 输出目录 {_rel_or_abs(output_root)}")
print(f" 配置文件 {config.get('_path')}")
print(line)
print(" 按 Ctrl+C 退出")
print(line, flush=True)
def main(argv: Optional[List[str]] = None) -> int:
args = parse_args(argv)
preflight()
config = merge_config(args)
try:
server, registry, models, node_dir, model_dirs, output_root = build_app(config)
except Exception as exc:
print(f"[启动失败] {type(exc).__name__}: {exc}", file=sys.stderr)
return 1
# 首启或显式要求时落盘配置(补全探测到的路径,便于下次直接改文件)
config_to_save = {k: v for k, v in config.items() if not k.startswith("_")}
config_to_save["node_dir"] = str(node_dir)
if config.get("_first_run") or args.save_config:
save_config(Path(str(config["_path"])), config_to_save)
if args.check:
print_banner(server, models, node_dir, model_dirs, output_root, config)
print("自检完成(--check 模式,不启动服务)")
return 0
print_banner(server, models, node_dir, model_dirs, output_root, config)
if config.get("open_browser"):
host, port = server.server_address[0], server.server_address[1]
shown = "127.0.0.1" if host in ("0.0.0.0", "::") else host
threading.Thread(
target=lambda: (time.sleep(1.0), webbrowser.open(f"http://{shown}:{port}/")),
daemon=True,
).start()
try:
server.serve_forever(poll_interval=0.4)
except KeyboardInterrupt:
print("\n收到 Ctrl+C,正在退出…")
finally:
server.engine.unload()
server.server_close()
return 0
if __name__ == "__main__":
sys.exit(main())