Files
BiRefNet_WebUI/selftest.py
T
2026-10-08 09:46:47 +08:00

180 lines
7.5 KiB
Python
Raw 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 -*-
"""端到端自检脚本:验证依赖、模型加载与抠图结果。
用法(本项目自包含,直接用自带的 Python 运行时):
python\\python.exe selftest.py # 用内置合成图测试第一个可用模型
python\\python.exe selftest.py --image photo.jpg # 用真实图片测试
python\\python.exe selftest.py --device cpu # 强制 CPU
python\\python.exe selftest.py --thorough # 把所有模型都跑一遍(慢)
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parent
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
import numpy as np # noqa: E402
from PIL import Image, ImageDraw # noqa: E402
def make_test_image(width: int = 900, height: int = 640) -> Image.Image:
"""合成一张「纯色背景 + 高对比主体」的测试图,便于断言遮罩是否合理。"""
y, x = np.mgrid[0:height, 0:width].astype(np.float32)
bg = np.stack(
[120 + 90 * x / width, 130 + 60 * y / height, 200 - 90 * x / width], axis=-1
).astype(np.uint8)
image = Image.fromarray(bg, "RGB")
draw = ImageDraw.Draw(image)
# 主体:一个亮色椭圆(占画面约 20% 面积)
box = (width * 0.25, height * 0.18, width * 0.75, height * 0.82)
draw.ellipse(box, fill=(245, 240, 235))
draw.ellipse((width * 0.36, height * 0.30, width * 0.42, height * 0.40), fill=(40, 40, 45))
draw.ellipse((width * 0.58, height * 0.30, width * 0.64, height * 0.40), fill=(40, 40, 45))
return image
def stats(arr: np.ndarray, name: str) -> None:
print(
f" {name:<12} 形状={arr.shape} 最小值={arr.min():.3f} "
f"最大值={arr.max():.3f} 均值={arr.mean():.3f}"
)
def main(argv=None) -> int:
parser = argparse.ArgumentParser(description="BiRefNet WebUI 自检")
parser.add_argument("--image", default=None, help="测试图片路径(默认使用内置合成图)")
parser.add_argument("--model", default=None, help="模型 key / 名称(默认取第一个可用模型)")
parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"])
parser.add_argument("--dtype", default="auto", choices=["auto", "float32", "float16", "bfloat16"])
parser.add_argument("--thorough", action="store_true", help="遍历所有模型")
parser.add_argument("--final-side", type=int, default=0,
help="最终尺寸·最长边(0 = 保持原图尺寸;>0 = 等比缩放到该最长边)")
parser.add_argument("--out", default=str(PROJECT_ROOT / "outputs" / "selftest"), help="结果输出目录")
args = parser.parse_args(argv)
from birefnet_web import compat
from birefnet_web.engine import InferenceEngine, ModelRegistry, describe_environment
from birefnet_web import imageops
print("=" * 68)
print("BiRefNet WebUI 自检")
print("=" * 68)
node_dir = compat.find_node_dir(None)
print(f"模型代码 {node_dir}")
model_dirs = compat.default_model_dirs(node_dir)
registry = ModelRegistry(model_dirs)
models = registry.scan()
print(f"模型目录 {', '.join(str(d) for d in model_dirs)}")
for m in models:
print(f" · {m.name:<22} {m.size_mb:>7.1f} MB arch={m.arch} bb={m.backbone} {m.path}")
if not models:
print("[失败] 没有找到任何模型权重,请先准备 *.safetensors")
return 1
env = describe_environment(InferenceEngine(registry, node_dir))
print(f"运行环境 python {env['python']} · torch {env['torch']} · cuda={env['cuda']}")
for dev in env["devices"]:
print(f" [{dev['index']}] {dev['name']} {dev['free_mem']}G/{dev['total_mem']}G")
if args.image:
src = Image.open(args.image)
src.load()
print(f"测试图片 {args.image} {src.size}")
else:
src = make_test_image()
print(f"测试图片 内置合成图 {src.size}")
engine = InferenceEngine(registry, node_dir=node_dir, max_cached=1)
targets = models if args.thorough else [registry.get(args.model) if args.model else models[0]]
out_dir = Path(args.out)
out_dir.mkdir(parents=True, exist_ok=True)
failures = 0
for info in targets:
print("-" * 68)
print(f"模型 {info.name}(arch={info.arch}, bb={info.backbone})")
stages = []
t0 = time.perf_counter()
try:
result = engine.remove_background(
src,
{
"model": info.key,
"device": args.device,
"dtype": args.dtype,
"resolution_mode": "square",
"width": 1024,
"height": 1024,
"upscale_method": "bilinear",
"mask_threshold": 0.0,
"refine_foreground": True,
"background": "transparent",
"output_mask": True,
"final_longest_side": args.final_side,
},
progress=lambda stage, frac: stages.append(f"{stage}({frac:.0%})"),
)
except Exception as exc:
failures += 1
print(f" [失败] {type(exc).__name__}: {exc}")
continue
cutout, mask = result["cutout"], result["mask"]
alpha = np.asarray(cutout.split()[-1]) if cutout.mode == "RGBA" else None
print(f" 设备/精度 {result['device']} / {result['dtype']} 输入 {result['input_size']}")
print(f" 输出 {cutout.mode} {cutout.size},耗时 {result['elapsed']}s")
print(f" 前景占比 {result['coverage']:.1%}")
print(f" 阶段 {' → '.join(stages)}")
if result.get("warning"):
print(f" [告警] {result['warning']}")
if alpha is not None:
stats(alpha.astype(np.float32) / 255.0, "alpha")
# 最终尺寸:结果尺寸必须等于「原图按比例缩放到最长边 = final_side」
expect_size = imageops.fit_longest_side((src.width, src.height), args.final_side)
if cutout.size != expect_size:
failures += 1
print(f" [失败] 最终尺寸不符:期望 {expect_size},实际 {cutout.size}")
else:
print(f" [通过] 最终尺寸 {cutout.size}(最长边参数 {args.final_side})")
cutout_path = out_dir / f"{info.name}_cutout.png"
imageops.save_image(cutout, str(cutout_path))
if mask is not None:
imageops.save_image(mask, str(out_dir / f"{info.name}_mask.png"))
print(f" 已保存 {cutout_path}")
if args.image is None:
coverage = result["coverage"]
if not 0.05 < coverage < 0.95:
failures += 1
print(f" [失败] 合成图期望前景占比在 5%~95% 之间,实际 {coverage:.1%}")
elif alpha is not None and not (
alpha.min() < 40 and alpha.max() > 215
):
failures += 1
print(" [失败] alpha 动态范围不足,遮罩可能异常")
else:
print(" [通过] 遮罩分布合理")
print(f" 总耗时 {time.perf_counter() - t0:.1f}s")
print("=" * 68)
if failures:
print(f"自检未通过:{failures} 项失败")
return 1
print("自检全部通过 ✔")
return 0
if __name__ == "__main__":
sys.exit(main())