feat: initial project setup
This commit is contained in:
@@ -0,0 +1,357 @@
|
||||
# -*- 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())
|
||||
Reference in New Issue
Block a user