180 lines
7.5 KiB
Python
180 lines
7.5 KiB
Python
# -*- 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())
|