File size: 3,476 Bytes
4cba0b8
 
 
 
fdb87aa
4cba0b8
 
 
 
fdb87aa
4cba0b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fdb87aa
 
4cba0b8
 
 
 
 
 
 
fdb87aa
 
 
4cba0b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
#!/usr/bin/env python3
"""Compare WD single vs batched ORT throughput on local images.

Usage (from backend/):
  ../.venv/Scripts/python.exe scripts/bench_batch_throughput.py --root /path/to/images --n 24
"""
from __future__ import annotations

import argparse
import os
import time
from pathlib import Path

IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}


def _collect_images(root: Path, limit: int) -> list[Path]:
    paths: list[Path] = []
    for path in root.rglob("*"):
        if path.is_file() and path.suffix.lower() in IMAGE_EXTS:
            paths.append(path)
            if len(paths) >= limit:
                break
    return paths


def _bench_batches(images: list[Path], batch_size: int, tagger_model: str, thr: float) -> dict:
    from app.services import extract_scores, extract_scores_batch

    # Warmup
    extract_scores(images[0], tagger_model=tagger_model, wd_general_threshold=thr)

    start = time.perf_counter()
    if batch_size <= 1:
        for image in images:
            extract_scores(image, tagger_model=tagger_model, wd_general_threshold=thr)
    else:
        for idx in range(0, len(images), batch_size):
            chunk = images[idx : idx + batch_size]
            extract_scores_batch(
                chunk,
                tagger_model=tagger_model,
                wd_general_threshold=thr,
            )
    elapsed = time.perf_counter() - start
    n = len(images)
    return {
        "batch_size": batch_size,
        "images": n,
        "elapsed_s": elapsed,
        "ms_per_image": (elapsed / n) * 1000.0,
        "images_per_min": (n / elapsed) * 60.0 if elapsed else 0.0,
    }


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--root",
        default=os.environ.get("THR3SHR_BENCH_ROOT", ""),
        help="Image root (or set THR3SHR_BENCH_ROOT)",
    )
    parser.add_argument("--n", type=int, default=24)
    parser.add_argument("--model", default="wd_swinv2_v3")
    parser.add_argument("--threshold", type=float, default=0.35)
    parser.add_argument("--batches", default="1,4,8")
    args = parser.parse_args()

    if not str(args.root).strip():
        raise SystemExit("Pass --root /path/to/images or set THR3SHR_BENCH_ROOT")

    root = Path(args.root)
    images = _collect_images(root, args.n)
    if len(images) < 4:
        raise SystemExit(f"Need at least 4 images under {root}, found {len(images)}")

    print(f"root={root}")
    print(f"model={args.model} n={len(images)} thr={args.threshold}")
    print(f"sample={[p.name for p in images[:3]]}")

    results = []
    for raw in args.batches.split(","):
        batch_size = int(raw.strip())
        print(f"\n=== batch_size={batch_size} ===", flush=True)
        row = _bench_batches(images, batch_size, args.model, args.threshold)
        results.append(row)
        print(
            f"  elapsed={row['elapsed_s']:.2f}s  "
            f"ms/img={row['ms_per_image']:.0f}  "
            f"img/min={row['images_per_min']:.1f}",
            flush=True,
        )

    baseline = next((r for r in results if r["batch_size"] == 1), None)
    if baseline:
        print("\n=== vs batch=1 ===")
        for row in results:
            if row["batch_size"] == 1:
                continue
            speedup = baseline["ms_per_image"] / row["ms_per_image"]
            print(f"  batch={row['batch_size']}: {speedup:.2f}x faster")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())