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