instance-2 / docqa /bench /accuracy.py
validops-east-3's picture
Deploy de83ed8cb4e1bbfa6f888d334f247492ce1b5bbf
9cacfae verified
Raw History Blame Contribute Delete
4.27 kB
#!/usr/bin/env python3
"""Accuracy check for the docqa service against corpus ground truth.
python bench/accuracy.py --corpus <dir> [--limit N] [--base URL]
Every number printed is measured from a real run against a running server.
Nothing here is estimated.
"""
from __future__ import annotations
import argparse
import json
import statistics
import subprocess
import time
from pathlib import Path
# Requested key -> ground-truth field in corpus/ground_truth.json
KEY_MAP = [
("INVOICE NO", "invoice_number"),
("Invoice Date", "invoice_date"),
("Due Date", "due_date"),
("PO Number", "po_number"),
("Payment Term", "payment_terms"),
("Currency", "currency"),
("Vendor Name", "vendor_name"),
("Vendor Address", "vendor_address"),
("Vendor Id", "vendor_tax_id"),
("GSTIN", "gstin"),
("Bill To", "bill_to"),
("Subtotal", "subtotal"),
("Tax Amount", "tax_amount"),
("Total Amount", "total"),
("Tax Label", "tax_label"),
]
def extract(base: str, pdf: Path, keys: list[str]) -> tuple[dict, float]:
started = time.perf_counter()
proc = subprocess.run(
["curl.exe", "-sS", "-X", "POST", f"{base}/v1/extract",
"-F", f"file=@{pdf}", "-F", "keys=" + ",".join(keys)],
capture_output=True, text=True,
)
wall = (time.perf_counter() - started) * 1000
if proc.returncode != 0:
raise RuntimeError(proc.stderr.strip() or "curl failed")
return json.loads(proc.stdout), wall
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--base", default="http://127.0.0.1:7860")
ap.add_argument("--corpus", type=Path, required=True)
ap.add_argument("--limit", type=int, default=1)
args = ap.parse_args()
manifest = json.loads((args.corpus / "ground_truth.json").read_text())
records = manifest[: args.limit] if args.limit else manifest
keys = [k for k, _ in KEY_MAP]
print(f"target {args.base}")
print(f"model impira/layoutlm-invoices keys={len(keys)} "
f"documents={len(records)}\n")
per_key: dict[str, list[int]] = {k: [] for k in keys}
latencies: list[float] = []
for record in records:
body, wall = extract(args.base, Path(record["path"]), keys)
latencies.append(body.get("latency_ms", 0.0))
print(f"{record['document_id']} variant={record['variant']} "
f"source={body.get('source')} words={body.get('word_count')} "
f"pages={body.get('pages_processed')}")
print(f"server={body.get('latency_ms', 0):.0f} ms "
f"wall={wall:.0f} ms (includes model load on first call)\n")
by_key = {f["key"]: f for f in body.get("fields", [])}
print(f" {'KEY':<14} {'EXPECTED':<26} {'STATUS':<15} "
f"{'GOT':<28} CONF")
print(" " + "-" * 92)
for key, truth in KEY_MAP:
expected = str(record[truth])
field = by_key.get(key, {})
got = (field.get("value") or "").strip()
ok = got == expected.strip()
per_key[key].append(1 if ok else 0)
print(f" {key:<14} {expected[:26]:<26} "
f"{str(field.get('status')):<15} "
f"{(got[:26] + (' OK' if ok else ' BAD')):<28} "
f"{field.get('confidence', 0):.4f}")
print(" " + "-" * 92 + "\n")
total_hit = sum(sum(v) for v in per_key.values())
total_n = sum(len(v) for v in per_key.values())
print("PER-KEY EXACT MATCH")
print("-" * 46)
for key, _ in KEY_MAP:
scores = per_key[key]
if scores:
acc = 100 * sum(scores) / len(scores)
print(f" {key:<16} {acc:6.1f}% ({sum(scores)}/{len(scores)})")
if total_n:
print(f"\n {'OVERALL':<16} {100*total_hit/total_n:6.1f}% "
f"({total_hit}/{total_n})")
if latencies:
print("\nLATENCY (whole document, all keys)")
print("-" * 46)
print(f" mean {statistics.mean(latencies):8.0f} ms")
print(f" min {min(latencies):8.0f} ms")
print(f" max {max(latencies):8.0f} ms")
print(f" per key {statistics.mean(latencies)/len(keys):6.0f} ms")
return 0
if __name__ == "__main__":
raise SystemExit(main())