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

387 lines
17 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 -*-
"""批量抠图自检:路径校验 / 命名规则 / 任务流水线 / HTTP 接口契约。
**不加载任何真实权重**(用假引擎直接产出结果图),所以跑起来只要几秒。
用法:
python\\python.exe selftest_batch.py # 全量
python selftest_batch.py --keep # 保留临时目录,便于看产物
python selftest_batch.py --logic-only # 只跑不依赖 torch 的纯逻辑部分
"""
from __future__ import annotations
import argparse
import json
import shutil
import sys
import tempfile
import threading
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import List, Optional
PROJECT_ROOT = Path(__file__).resolve().parent
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
_FAILS: List[str] = []
def check(name: str, cond: bool, extra: object = "") -> None:
print(f" {'[通过]' if cond else '[失败]'} {name}" + ("" if cond else f" ← {extra}"))
if not cond:
_FAILS.append(name)
# --------------------------------------------------------------------------- #
# 1) 纯逻辑:路径校验 + 命名规则(不需要 torch)
# --------------------------------------------------------------------------- #
def part_logic(tmp: Path) -> None:
from birefnet_web import batch
src, out = tmp / "src", tmp / "out"
out.mkdir(parents=True)
(src / "sub").mkdir(parents=True)
(out / "RMBG_already_1.png").write_bytes(b"x")
(src / "photo.jpg").write_bytes(b"x")
(src / "a_pretty_long_original_filename_here.png").write_bytes(b"x")
(src / "note.txt").write_bytes(b"x")
(src / "RMBG_old_1.png").write_bytes(b"x")
(src / "sub" / "deep.webp").write_bytes(b"x")
print("\n[1] 路径校验")
missing = tmp / "nope"
for label, fn in (("输入", batch.check_input_dir), ("输出", batch.check_output_dir)):
try:
fn(str(missing), PROJECT_ROOT)
check(f"{label}目录不存在时抛错", False, "没有抛 BatchPathError")
except batch.BatchPathError as exc:
check(f"{label}目录不存在时抛错", True)
check(f"{label}报错时未创建目录", not missing.exists())
print(f" → {exc}")
try:
batch.check_input_dir(str(src / "photo.jpg"), PROJECT_ROOT)
check("输入是文件时抛错", False, "没有抛 BatchPathError")
except batch.BatchPathError as exc:
check("输入是文件时抛错", True)
print(f" → {exc}")
try:
batch.check_output_dir(str(PROJECT_ROOT), PROJECT_ROOT, allow_system_dir=False)
check("项目目录可写时通过", True)
except batch.BatchPathError as exc:
check("项目目录可写时通过", False, exc)
check("相对路径基于项目根解析", batch.resolve_user_path("outputs", PROJECT_ROOT) == PROJECT_ROOT / "outputs")
check("成对引号被去掉", batch.resolve_user_path(' "outputs" ', PROJECT_ROOT) == PROJECT_ROOT / "outputs")
print("\n[2] 目录扫描")
names = sorted(p.name for p in batch.scan_images(src))
check("非递归只扫顶层图片", names == ["a_pretty_long_original_filename_here.png", "photo.jpg"], names)
deep = batch.scan_images(src, recursive=True)
check("递归包含子目录", len(deep) == 3, [p.name for p in deep])
check("跳过非图片文件", all(p.suffix != ".txt" for p in deep))
check("跳过本工具产物 RMBG_*", all(not p.name.startswith("RMBG_") for p in deep))
check("跳过输出目录", all(out not in p.parents for p in batch.scan_images(src, recursive=True, skip_dirs=[out])))
print("\n[3] 命名规则")
long_src = src / "a_pretty_long_original_filename_here.png"
check("RMBG_ + 原主干 + _ + 时间戳 + 原后缀",
batch.build_output_name(src / "photo.jpg", 1700000000, "color") == "RMBG_photo_1700000000.jpg")
check("主干超 20 字符截断",
batch.build_output_name(long_src, 1700000000, "transparent") == "RMBG_a_pretty_long_origin_1700000000.png",
batch.build_output_name(long_src, 1700000000, "transparent"))
check("中文按字符截断(不是按字节)",
batch.build_output_name(Path("一二三四五六七八九十一二三四五六七八九十一二三.png"), 7, "transparent")
== "RMBG_一二三四五六七八九十一二三四五六七八九十_7.png")
check("透明 + jpg 退化 .png(jpg 装不下 alpha)",
batch.build_output_name(src / "photo.jpg", 1, "transparent") == "RMBG_photo_1.png")
check("纯色 + jpg 保持 .jpg",
batch.build_output_name(src / "photo.jpg", 1, "color") == "RMBG_photo_1.jpg")
check("非法字符被清洗", batch.build_output_name(Path("a?c*d.jpg"), 1, "color") == "RMBG_a_c_d_1.jpg")
check("同名文件避让为 _1", batch.unique_path(out, "RMBG_already_1.png").name == "RMBG_already_1_1.png")
check("无同名时原样返回", batch.unique_path(out, "RMBG_new_1.png").name == "RMBG_new_1.png")
print("\n[4] 最终尺寸 · 最长边换算")
from birefnet_web import fsutil
fit = fsutil.fit_longest_side
check("0 = 保持原图尺寸", fit((2000, 3000), 0) == (2000, 3000))
check("负数同样按不缩放处理", fit((2000, 3000), -1) == (2000, 3000))
check("需求例 2:竖图 2000×3000 → 1440 得 960×1440", fit((2000, 3000), 1440) == (960, 1440), fit((2000, 3000), 1440))
check("需求例 3:横图 3000×2000 → 1440 得 1440×960", fit((3000, 2000), 1440) == (1440, 960), fit((3000, 2000), 1440))
check("最长边已等于目标值时不动", fit((1440, 960), 1440) == (1440, 960))
check("小图按比例放大(800×600 → 1600 得 1600×1200)", fit((800, 600), 1600) == (1600, 1200), fit((800, 600), 1600))
check("正方形等比缩放", fit((1000, 1000), 500) == (500, 500))
check("极端比例也不会算出 0 像素", fit((10000, 3), 10) == (10, 1), fit((10000, 3), 10))
# --------------------------------------------------------------------------- #
# 2) 任务流水线(假引擎,不加载权重)
# --------------------------------------------------------------------------- #
class FakeModel:
key = "fake:safetensors"
name = "Fake"
def to_dict(self) -> dict:
return {"key": self.key, "name": self.name, "file": "fake.safetensors", "arch": "v1",
"size_mb": 1.0, "backbone": "swin_v1_tiny", "path": "fake"}
class FakeRegistry:
model_dirs: List[Path] = []
def list(self) -> List[FakeModel]:
return [FakeModel()]
def get(self, key: str) -> FakeModel:
if key != FakeModel.key:
raise KeyError(key)
return FakeModel()
class FakeEngine:
"""只按参数产出一张图,用来验证落盘与命名,不碰 torch。"""
def __init__(self) -> None:
self.calls = 0
def unload(self) -> None:
pass
def remove_background(self, image, options, progress=None): # noqa: ANN001
self.calls += 1
if progress:
progress("推理", 0.5)
size = image.size
transparent = options.get("background") == "transparent"
rgba = image.convert("RGBA")
rgba.putalpha(128)
cutout = rgba if transparent else rgba.convert("RGB")
return {
"cutout": cutout,
"mask": image.convert("L") if options.get("output_mask") else None,
"width": size[0], "height": size[1], "coverage": 0.5, "elapsed": 0.01,
"input_size": (1024, 1024), "warning": None, "device": "cpu", "dtype": "float32",
}
def _seed_images(root: Path) -> None:
from PIL import Image
root.mkdir(parents=True, exist_ok=True)
Image.new("RGB", (64, 48), (200, 30, 30)).save(root / "photo.jpg")
Image.new("RGB", (32, 32), (30, 200, 30)).save(root / "a_pretty_long_original_filename_here.png")
(root / "note.txt").write_text("not an image", encoding="utf-8")
def _wait(task, timeout: float = 30.0) -> bool:
deadline = time.time() + timeout
while time.time() < deadline:
if task.finished_at is not None:
return True
time.sleep(0.05)
return False
def part_pipeline(tmp: Path) -> None:
from birefnet_web.server import TaskManager
print("\n[4] 批量任务流水线(假引擎)")
src, out = tmp / "pipe_src", tmp / "pipe_out"
_seed_images(src)
# 两个不同子目录里的同名文件 → 检验同一秒内的重名避让
(src / "a").mkdir()
(src / "b").mkdir()
from PIL import Image
Image.new("RGB", (16, 16), (0, 0, 255)).save(src / "a" / "pic.png")
Image.new("RGB", (16, 16), (255, 0, 255)).save(src / "b" / "pic.png")
out.mkdir()
engine = FakeEngine()
manager = TaskManager(engine, tmp / "default_out", max_tasks=5)
options = {"background": "transparent", "output_mask": False}
task = manager.submit_batch(src, out, options, recursive=True)
check("任务模式标记为 batch", task.mode == "batch")
check("原图未被拷贝进任务目录(就地读取)",
all(Path(j.input_path).parent != tmp / "default_out" for j in task.images))
check("递归扫描命中 4 张", len(task.images) == 4, len(task.images))
check("等待处理结束", _wait(task))
check("任务状态为 done", task.state == "done", task.state)
produced = sorted(p.name for p in out.iterdir())
print(f" 产物:{produced}")
check("产物数量 = 图片数量", len(produced) == 4, produced)
check("全部带 RMBG_ 前缀", all(n.startswith("RMBG_") for n in produced))
check("jpg + 透明 → .png", any(n.startswith("RMBG_photo_") and n.endswith(".png") for n in produced))
check("长主干截断到 20 字符",
any(n.startswith("RMBG_a_pretty_long_origin_") for n in produced))
check("同名文件避让(无覆盖)",
sum(1 for n in produced if n.startswith("RMBG_pic_")) == 2, produced)
check("每张都记录了 output_name",
all(j.output_name and (out / j.output_name).is_file() for j in task.images))
# 纯色背景 + 同时输出遮罩
out2 = tmp / "pipe_out_color"
out2.mkdir()
task2 = manager.submit_batch(src / "a", out2, {"background": "color", "output_mask": True})
check("子目录只有 1 张", len(task2.images) == 1)
check("等待第二个任务结束", _wait(task2))
names2 = sorted(p.name for p in out2.iterdir())
print(f" 产物:{names2}")
check("纯色 + png 保持 .png", any(n.startswith("RMBG_pic_") and n.endswith(".png") and "_mask" not in n for n in names2))
check("同时输出遮罩 _mask.png", any(n.endswith("_mask.png") for n in names2), names2)
# 输入目录不存在 → 由 HTTP 层拦下,这里只验目录空的情况
try:
manager.submit_batch(tmp / "empty_in", out, options)
check("空目录抛 ValueError", False, "没有抛")
except ValueError as exc:
check("空目录抛 ValueError", True)
print(f" → {exc}")
# --------------------------------------------------------------------------- #
# 3) HTTP 契约(前端就是照这个调的)
# --------------------------------------------------------------------------- #
def _post(url: str, payload: dict) -> tuple:
body = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
try:
with urllib.request.urlopen(req, timeout=10) as res: # noqa: S310
return res.status, json.loads(res.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
return exc.code, json.loads(exc.read().decode("utf-8"))
def _get(url: str) -> tuple:
try:
with urllib.request.urlopen(url, timeout=10) as res: # noqa: S310
return res.status, json.loads(res.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
return exc.code, json.loads(exc.read().decode("utf-8"))
def part_http(tmp: Path) -> None:
from birefnet_web.server import TaskManager, WebUIServer
print("\n[5] HTTP 接口 /api/batch")
engine = FakeEngine()
out_root = tmp / "http_default_out"
tasks = TaskManager(engine, out_root, max_tasks=5)
server = WebUIServer(
("127.0.0.1", 0),
registry=FakeRegistry(), # type: ignore[arg-type]
engine=engine, # type: ignore[arg-type]
tasks=tasks,
web_dir=PROJECT_ROOT / "web",
output_root=out_root,
node_dir=PROJECT_ROOT / "vendor",
logger=lambda msg: None,
project_root=PROJECT_ROOT,
)
port = server.server_address[1]
threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True).start()
base = f"http://127.0.0.1:{port}"
try:
src = tmp / "http_src"
_seed_images(src)
out = tmp / "http_out"
out.mkdir()
status, data = _post(f"{base}/api/batch", {"input_dir": str(tmp / "missing")})
check("输入目录不存在 → 400", status == 400, (status, data))
check("错误信息直接点出不存在", "不存在" in str(data.get("error")), data)
check("仍然没有创建该目录", not (tmp / "missing").exists())
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(tmp / "missing_out")})
check("输出目录不存在 → 400", status == 400, (status, data))
check("仍然没有创建输出目录", not (tmp / "missing_out").exists())
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(src)})
check("输入 == 输出时不会「一张都扫不到」", status == 200 and data.get("count") == 2, (status, data))
if status == 200:
_wait(tasks.get(data["task_id"]))
for f in sorted(src.glob("RMBG_*")):
f.unlink()
status, data = _post(f"{base}/api/batch",
{"input_dir": str(src), "output_dir": str(out), "options": {"background": "transparent"}})
check("正常提交 → 200", status == 200 and data.get("count") == 2, (status, data))
task_id = data.get("task_id")
check("回传输出目录绝对路径", data.get("output_dir") == str(out), data.get("output_dir"))
if task_id:
_wait(tasks.get(task_id))
status, detail = _get(f"{base}/api/tasks/{task_id}?tail=1")
task = detail.get("task", {})
check("任务详情带 mode/input_dir/output_dir",
task.get("mode") == "batch" and task.get("input_dir") == str(src) and task.get("output_dir") == str(out),
task.get("mode"))
check("tail=1 只回传 1 张(载荷可控)", len(task.get("images", [])) == 1, len(task.get("images", [])))
check("images_total = 2", task.get("images_total") == 2, task.get("images_total"))
check("images_from 指向窗口起点", task.get("images_from") == 1, task.get("images_from"))
check("批量任务未被塞进单图结果网格的判据(mode=batch)", task.get("mode") == "batch")
opts = task.get("options") or {}
check("未指定的参数落到新默认值(不输出遮罩 / 最终尺寸不限)",
opts.get("output_mask") is False and opts.get("final_longest_side") == 0,
{k: opts.get(k) for k in ("output_mask", "final_longest_side")})
status, listing = _get(f"{base}/api/tasks")
modes = [t.get("mode") for t in listing.get("tasks", [])]
check("任务列表里能区分两种模式", set(modes) <= {"batch", "upload"} and "batch" in modes, modes)
status, bad = _post(f"{base}/api/batch", {"input_dir": str(src), "recursive": True})
check("留空输出目录 → 落到项目 outputs", status == 200 and bad.get("output_dir") == str(out_root),
bad.get("output_dir"))
if status == 200:
_wait(tasks.get(bad["task_id"]))
finally:
server.shutdown()
server.server_close()
# --------------------------------------------------------------------------- #
def main(argv: Optional[List[str]] = None) -> int:
parser = argparse.ArgumentParser(description="批量抠图自检")
parser.add_argument("--keep", action="store_true", help="保留临时目录")
parser.add_argument("--logic-only", action="store_true", help="只跑不依赖 torch 的纯逻辑部分")
args = parser.parse_args(argv)
tmp = Path(tempfile.mkdtemp(prefix="birefnet_batch_"))
print("=" * 68)
print("批量抠图自检")
print(f"临时目录 {tmp}")
print("=" * 68)
try:
part_logic(tmp)
if args.logic_only:
print("\n[跳过] 任务流水线与 HTTP 契约(--logic-only)")
else:
try:
part_pipeline(tmp)
part_http(tmp)
except ImportError as exc:
print(f"\n[跳过] 任务流水线与 HTTP 契约(缺少依赖:{exc})")
finally:
if args.keep:
print(f"\n临时目录已保留:{tmp}")
else:
shutil.rmtree(tmp, ignore_errors=True)
print("=" * 68)
if _FAILS:
print(f"自检未通过:{len(_FAILS)} 项失败")
for name in _FAILS:
print(f" · {name}")
return 1
print("自检全部通过 ✔")
return 0
if __name__ == "__main__":
sys.exit(main())