Add Accuracy Benchmark for OCR models (#25364)
Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
co-authored by
Ma Mingfei
parent
e2b55bdbab
commit
b3ab56545b
@@ -0,0 +1,171 @@
|
|||||||
|
# OCR Accuracy Benchmark
|
||||||
|
|
||||||
|
Evaluates `deepseek-ai/DeepSeek-OCR-2` (and any compatible OCR VLM) on
|
||||||
|
**olmOCR-bench** (AllenAI), the benchmark explicitly used in DeepSeek-OCR-2
|
||||||
|
official evaluations.
|
||||||
|
|
||||||
|
Targets **olmOCR-bench** because:
|
||||||
|
- Public HuggingFace dataset with 7,010 deterministic unit tests
|
||||||
|
- Explicitly cited by DeepSeek-OCR-2 authors
|
||||||
|
- Clear pass/fail semantics — no heavy CDM/TEDS/LaTeXML dependencies
|
||||||
|
- Covers 7 challenging document types across 1,403 PDF pages
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Step 0 (one-time): download olmOCR-bench including PDFs (~2 GB via Git LFS)
|
||||||
|
pip install huggingface_hub
|
||||||
|
hf download --repo-type dataset \
|
||||||
|
allenai/olmOCR-bench --local-dir ./olmOCR-bench
|
||||||
|
# This places bench_data/ (7 JSONL files + pdfs/ directory) under ./olmOCR-bench/
|
||||||
|
|
||||||
|
# Required: benchmark dependencies (pymupdf is in sglang[test]; aiohttp/tqdm are in core)
|
||||||
|
pip install "sglang[test]"
|
||||||
|
# OR install PDF rendering manually (choose one):
|
||||||
|
# pip install pymupdf # recommended (faster, pure Python wheel)
|
||||||
|
# pip install pdf2image # needs poppler: sudo apt install poppler-utils
|
||||||
|
|
||||||
|
# Start the sglang server (matches run.sh in this repo)
|
||||||
|
python -m sglang.launch_server \
|
||||||
|
--model-path deepseek-ai/DeepSeek-OCR-2 \
|
||||||
|
--host 127.0.0.1 --port 30000
|
||||||
|
```
|
||||||
|
|
||||||
|
> **Why the download step?**
|
||||||
|
> The olmOCR-bench PDF files are stored in Git LFS on HuggingFace.
|
||||||
|
> `datasets.load_dataset()` cannot retrieve LFS-backed binary files, so the
|
||||||
|
> benchmark reads the JSONL test files and PDFs directly from a local clone of
|
||||||
|
> the repository.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Full benchmark — all 7 splits (~7,010 tests)
|
||||||
|
python -m benchmark.ocr.bench_sglang \
|
||||||
|
--port 30000 \
|
||||||
|
--model deepseek-ai/DeepSeek-OCR-2 \
|
||||||
|
--split all \
|
||||||
|
--concurrency 8 \
|
||||||
|
--output-dir ./ocr_bench_results
|
||||||
|
|
||||||
|
# Single split
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split arxiv_math --concurrency 16
|
||||||
|
|
||||||
|
# Quick smoke-test (50 samples from one split)
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split old_scans --max-samples 50
|
||||||
|
|
||||||
|
# Use "Free OCR" prompt instead of markdown conversion
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split all --prompt-mode free_ocr
|
||||||
|
|
||||||
|
# Save raw model outputs for inspection
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split multi_column --save-raw-outputs
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Arguments
|
||||||
|
|
||||||
|
| Argument | Default | Description |
|
||||||
|
|----------|---------|-------------|
|
||||||
|
| `--port` | `30000` | sglang server port |
|
||||||
|
| `--host` | `127.0.0.1` | sglang server host |
|
||||||
|
| `--model` | `deepseek-ai/DeepSeek-OCR-2` | Model ID (must match running server) |
|
||||||
|
| `--split` | `all` | Split name or `all` |
|
||||||
|
| `--concurrency` | `8` | Concurrent requests to server |
|
||||||
|
| `--output-dir` | `./ocr_bench_results` | Directory for result JSON files |
|
||||||
|
| `--max-samples` | `-1` | Limit samples per split (-1 = all) |
|
||||||
|
| `--prompt-mode` | `markdown` | `markdown` or `free_ocr` |
|
||||||
|
| `--request-timeout` | `300` | Per-request timeout (seconds) |
|
||||||
|
| `--render-dpi` | `150` | DPI for PDF → PNG rendering |
|
||||||
|
| `--save-raw-outputs` | `False` | Include raw OCR text in JSON output |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Test Classes (olmOCR-bench)
|
||||||
|
|
||||||
|
| Test Type | Description | Matching strategy |
|
||||||
|
|-----------|-------------|-------------------|
|
||||||
|
| `text_presence` | 1–3 sentence text must appear in OCR output | Exact or fuzzy; optional position constraint (first/last N chars) |
|
||||||
|
| `text_absence` | Header/footer/page-number text must NOT appear | Fuzzy; case-insensitive |
|
||||||
|
| `natural_reading_order` | Two text spans must appear in the correct order | Soft/fuzzy positional matching |
|
||||||
|
| `table_accuracy` | Cell value with correct neighbor relationship | Markdown + HTML table parsing |
|
||||||
|
| `math_formula_accuracy` | LaTeX key-token symbols present in math regions | Symbol-token matching (≥70% threshold) |
|
||||||
|
|
||||||
|
> **Note on math**: The official olmOCR-bench uses KaTeX rendering + Playwright for
|
||||||
|
> bounding-box symbol matching. This benchmark uses a symbol-token proxy (no browser
|
||||||
|
> dependency). Scores on `arxiv_math` and `old_scans_math` may therefore differ from
|
||||||
|
> the official leaderboard.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Dataset Splits
|
||||||
|
|
||||||
|
| Split | Documents | Tests | Document type |
|
||||||
|
|-------|-----------|-------|---------------|
|
||||||
|
| `arxiv_math` | 522 | 2,927 | arXiv math papers |
|
||||||
|
| `old_scans_math` | 36 | 458 | Scanned math textbooks (Internet Archive) |
|
||||||
|
| `table_tests` | 188 | 1,020 | Documents with tables |
|
||||||
|
| `old_scans` | 98 | 526 | Historical / typewritten documents (Library of Congress) |
|
||||||
|
| `headers_footers` | 266 | 753 | Documents with headers/footers to exclude |
|
||||||
|
| `multi_column` | 231 | 884 | Multi-column layouts |
|
||||||
|
| `long_tiny_text` | 62 | 442 | Dense small-print pages |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Reference Scores
|
||||||
|
|
||||||
|
Column order matches the [olmOCR README](https://github.com/allenai/olmocr): AR = arxiv_math, OSM = old_scans_math, TA = table_tests, OS = old_scans, HF = headers_footers, MC = multi_column, LTT = long_tiny_text, Base = baseline.
|
||||||
|
|
||||||
|
| Model | AR | OSM | TA | OS | HF | MC | LTT | Base | **Overall** |
|
||||||
|
|-------|:--:|:---:|:--:|:--:|:--:|:--:|:---:|:----:|:-----------:|
|
||||||
|
| DeepSeek-OCR v1 | 77.2 | 73.6 | 80.2 | 33.3 | 96.1 | 66.4 | 79.4 | 99.8 | **75.7** |
|
||||||
|
| **DeepSeek-OCR-2** | **82.0** | **72.0** | **77.4** | — | — | — | — | — | **76.3** |
|
||||||
|
| olmOCR v0.4.0 | 83.0 | 82.3 | 84.9 | 47.7 | 96.1 | 83.7 | 81.9 | 99.7 | **82.4** |
|
||||||
|
| PaddleOCR-VL\* | 85.7 | 71.0 | 84.1 | 37.8 | 97.0 | 79.9 | 85.7 | 98.5 | **80.0** |
|
||||||
|
| Mistral OCR API | 77.2 | 67.5 | 60.6 | 29.3 | 93.6 | 71.3 | 77.1 | 99.4 | **72.0** |
|
||||||
|
| Marker 1.10.1 | 83.8 | 66.8 | 72.9 | 33.5 | 86.6 | 80.0 | 85.7 | 99.3 | **76.1** |
|
||||||
|
| MinerU 2.5.4\* | 76.6 | 54.6 | 84.9 | 33.7 | 96.6 | 78.2 | 83.5 | 93.7 | **75.2** |
|
||||||
|
|
||||||
|
\* = scores reported by model authors, not reproduced by olmOCR team.
|
||||||
|
|
||||||
|
DeepSeek-OCR-2 per-split scores for OS/HF/MC/LTT are not officially reported; only the three highlighted splits and overall appear on the [HuggingFace model card](https://huggingface.co/deepseek-ai/DeepSeek-OCR-2).
|
||||||
|
|
||||||
|
> **Note on math scores**: This benchmark uses token-overlap matching (≥70% threshold) rather than the official KaTeX rendering + Playwright bounding-box comparison. Scores on `arxiv_math` and `old_scans_math` will therefore differ from the official leaderboard.
|
||||||
|
|
||||||
|
Sources: [olmOCR README](https://github.com/allenai/olmocr), [DeepSeek-OCR-2 HF card](https://huggingface.co/deepseek-ai/DeepSeek-OCR-2).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Output Files
|
||||||
|
|
||||||
|
Results are written to `--output-dir`:
|
||||||
|
|
||||||
|
```
|
||||||
|
ocr_bench_results/
|
||||||
|
├── arxiv_math.json # per-split detailed results
|
||||||
|
├── old_scans.json
|
||||||
|
├── ...
|
||||||
|
└── summary.json # aggregated across all evaluated splits
|
||||||
|
```
|
||||||
|
|
||||||
|
Each split JSON contains:
|
||||||
|
- `overall_score`: % tests passed
|
||||||
|
- `by_type`: per-test-type pass rate
|
||||||
|
- `total_tests`, `total_passed`, `error_samples`
|
||||||
|
- Per-sample `test_results` with `type`, `passed`, optional `error`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Files
|
||||||
|
|
||||||
|
| File | Description |
|
||||||
|
|------|-------------|
|
||||||
|
| `bench_sglang.py` | Main benchmark runner — loads dataset, sends requests, aggregates |
|
||||||
|
| `eval_utils.py` | Test evaluators, Normalized Edit Distance metric, aggregation helpers |
|
||||||
|
| `generate_report.py` | Generates self-contained HTML reports with MathJax from result JSONs |
|
||||||
|
| `README.md` | This file |
|
||||||
@@ -0,0 +1,727 @@
|
|||||||
|
"""
|
||||||
|
Benchmark DeepSeek-OCR-2 (and similar OCR VLMs) on olmOCR-bench via a running sglang server.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
# 0. Download the dataset (one-time, ~2 GB with PDFs via Git LFS)
|
||||||
|
hf download --repo-type dataset \\
|
||||||
|
allenai/olmOCR-bench --local-dir ./olmOCR-bench
|
||||||
|
|
||||||
|
# 1. Start the sglang server (matches run.sh)
|
||||||
|
python -m sglang.launch_server \\
|
||||||
|
--model-path deepseek-ai/DeepSeek-OCR-2 --host 127.0.0.1 --port 30000
|
||||||
|
|
||||||
|
# 2. Run the full benchmark (all 7 splits, ~7,010 tests)
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split all --concurrency 8
|
||||||
|
|
||||||
|
# 3. Quick run on a single split
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split arxiv_math --concurrency 16
|
||||||
|
|
||||||
|
# 4. Limit pages for a fast smoke-test
|
||||||
|
python -m benchmark.ocr.bench_sglang --port 30000 --split old_scans --max-samples 10
|
||||||
|
|
||||||
|
# 5. Custom dataset location
|
||||||
|
python -m benchmark.ocr.bench_sglang --bench-dir /data/olmOCR-bench/bench_data
|
||||||
|
|
||||||
|
Dataset:
|
||||||
|
allenai/olmOCR-bench – 7 splits, 1,403 PDFs, 7,010 unit tests
|
||||||
|
Splits: arxiv_math | old_scans_math | table_tests | old_scans |
|
||||||
|
headers_footers | multi_column | long_tiny_text
|
||||||
|
PDFs are stored via Git LFS; hf download (step 0) is required.
|
||||||
|
|
||||||
|
Reference scores (olmOCR-bench):
|
||||||
|
DeepSeek-OCR v1 : 75.7 ± 1.0
|
||||||
|
DeepSeek-OCR-2 : 76.3 (reported on HF model card)
|
||||||
|
olmOCR v0.4.0 : 82.4 ± 1.1
|
||||||
|
PaddleOCR-VL : 80.0 ± 1.0
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
import traceback
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Tuple
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
from tqdm.asyncio import tqdm as atqdm
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Paths: allow running from the repo root or from benchmark/ocr/
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
_SCRIPT_DIR = Path(__file__).resolve().parent
|
||||||
|
if str(_SCRIPT_DIR) not in sys.path:
|
||||||
|
sys.path.insert(0, str(_SCRIPT_DIR))
|
||||||
|
|
||||||
|
from eval_utils import (
|
||||||
|
aggregate_results,
|
||||||
|
evaluate_olmocr_tests,
|
||||||
|
print_results_table,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Constants
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
OLMOCR_BENCH_SPLITS = [
|
||||||
|
"arxiv_math",
|
||||||
|
"old_scans_math",
|
||||||
|
"table_tests",
|
||||||
|
"old_scans",
|
||||||
|
"headers_footers",
|
||||||
|
"multi_column",
|
||||||
|
"long_tiny_text",
|
||||||
|
]
|
||||||
|
|
||||||
|
# DeepSeek-OCR-2 prompt formats (https://github.com/deepseek-ai/DeepSeek-OCR-2)
|
||||||
|
_PROMPT_MARKDOWN = "<|grounding|>Convert the document to markdown."
|
||||||
|
_PROMPT_FREE_OCR = "Free OCR."
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Argument dataclass
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BenchArgs:
|
||||||
|
port: int = 30000
|
||||||
|
host: str = "127.0.0.1"
|
||||||
|
model: str = "deepseek-ai/DeepSeek-OCR-2"
|
||||||
|
split: str = "all"
|
||||||
|
concurrency: int = 8
|
||||||
|
output_dir: str = "./ocr_bench_results"
|
||||||
|
max_samples: int = -1
|
||||||
|
prompt_mode: str = "markdown"
|
||||||
|
bench_dir: str = "./olmOCR-bench/bench_data"
|
||||||
|
request_timeout: int = 300
|
||||||
|
save_raw_outputs: bool = False
|
||||||
|
render_dpi: int = 150
|
||||||
|
debug: bool = False
|
||||||
|
debug_accuracy: bool = False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def add_cli_args(parser: argparse.ArgumentParser) -> None:
|
||||||
|
parser.add_argument(
|
||||||
|
"--port", type=int, default=BenchArgs.port, help="sglang server port"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--host", type=str, default=BenchArgs.host, help="sglang server host"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--model",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.model,
|
||||||
|
help="Model identifier (must match the running server)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--split",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.split,
|
||||||
|
choices=OLMOCR_BENCH_SPLITS + ["all"],
|
||||||
|
help="Dataset split to evaluate. Use 'all' for all splits.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--concurrency",
|
||||||
|
type=int,
|
||||||
|
default=BenchArgs.concurrency,
|
||||||
|
help="Max concurrent requests to the sglang server",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output-dir",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.output_dir,
|
||||||
|
help="Directory for result JSON files",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max-samples",
|
||||||
|
type=int,
|
||||||
|
default=BenchArgs.max_samples,
|
||||||
|
help="Max samples per split (-1 = all)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--prompt-mode",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.prompt_mode,
|
||||||
|
choices=["markdown", "free_ocr"],
|
||||||
|
help=(
|
||||||
|
"Prompt mode for the OCR model: "
|
||||||
|
"'markdown' → '<|grounding|>Convert the document to markdown.'; "
|
||||||
|
"'free_ocr' → 'Free OCR.'"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--bench-dir",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.bench_dir,
|
||||||
|
help=(
|
||||||
|
"Local directory containing the olmOCR-bench bench_data/ folder "
|
||||||
|
"(JSONL files + pdfs/ sub-directory). Download first with: "
|
||||||
|
"hf download --repo-type dataset allenai/olmOCR-bench "
|
||||||
|
"--local-dir ./olmOCR-bench"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--request-timeout",
|
||||||
|
type=int,
|
||||||
|
default=BenchArgs.request_timeout,
|
||||||
|
help="Per-request timeout in seconds",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save-raw-outputs",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help="Include raw OCR text in result JSON (useful for debugging)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--render-dpi",
|
||||||
|
type=int,
|
||||||
|
default=BenchArgs.render_dpi,
|
||||||
|
help="DPI for rendering PDF pages to images",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--debug",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help=(
|
||||||
|
"Enable debug logging: print per-sample errors immediately, "
|
||||||
|
"show full tracebacks, and abort on the first server connection failure."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--debug-accuracy",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help=(
|
||||||
|
"Print per-sample accuracy details: input PDF path, expected "
|
||||||
|
"expressions/text, OCR output, and pass/fail per test."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_cli_args(cls, args: argparse.Namespace) -> "BenchArgs":
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
|
attrs = [f.name for f in dataclasses.fields(cls)]
|
||||||
|
return cls(**{attr: getattr(args, attr) for attr in attrs})
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Server preflight check
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def preflight_check(api_url: str, model: str, debug: bool) -> None:
|
||||||
|
"""
|
||||||
|
Send a minimal request to the server before the benchmark starts.
|
||||||
|
Raises SystemExit with a clear message if the server is unreachable or
|
||||||
|
returns an unexpected error.
|
||||||
|
"""
|
||||||
|
print(f" Preflight check → {api_url} … ", end="", flush=True)
|
||||||
|
payload = {
|
||||||
|
"model": model,
|
||||||
|
"messages": [{"role": "user", "content": "ping"}],
|
||||||
|
"max_tokens": 1,
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.post(
|
||||||
|
api_url,
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=10),
|
||||||
|
) as resp:
|
||||||
|
if resp.status in (200, 400): # 400 = bad request but server is alive
|
||||||
|
print("OK")
|
||||||
|
return
|
||||||
|
body = await resp.text()
|
||||||
|
print(f"FAILED (HTTP {resp.status})")
|
||||||
|
raise SystemExit(
|
||||||
|
f"Server returned HTTP {resp.status}.\nResponse: {body[:400]}\n"
|
||||||
|
f"Ensure the server is running: python -m sglang.launch_server "
|
||||||
|
f"--model-path {model} --host 127.0.0.1 --port <PORT>"
|
||||||
|
)
|
||||||
|
except (aiohttp.ClientConnectorError, asyncio.TimeoutError) as exc:
|
||||||
|
print("FAILED")
|
||||||
|
msg = (
|
||||||
|
f"Cannot reach sglang server at {api_url}\n"
|
||||||
|
f"Error: {exc}\n"
|
||||||
|
"Check:\n"
|
||||||
|
" 1. Is the server running? (python -m sglang.launch_server ...)\n"
|
||||||
|
" 2. Is --port correct?\n"
|
||||||
|
" 3. Are you running inside the same docker container as the server?"
|
||||||
|
)
|
||||||
|
if debug:
|
||||||
|
msg += f"\n\nFull traceback:\n{traceback.format_exc()}"
|
||||||
|
raise SystemExit(msg)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# PDF → base64 PNG
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def pdf_page_to_base64_png(pdf_bytes: bytes, page_num: int = 0, dpi: int = 150) -> str:
|
||||||
|
"""
|
||||||
|
Render a single PDF page to a base64-encoded PNG string.
|
||||||
|
|
||||||
|
Tries PyMuPDF (fitz) first; falls back to pdf2image / poppler.
|
||||||
|
page_num is 0-indexed.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
import fitz # PyMuPDF
|
||||||
|
|
||||||
|
doc = fitz.open(stream=pdf_bytes, filetype="pdf")
|
||||||
|
if page_num >= len(doc):
|
||||||
|
page_num = len(doc) - 1
|
||||||
|
page = doc[page_num]
|
||||||
|
mat = fitz.Matrix(dpi / 72.0, dpi / 72.0)
|
||||||
|
pix = page.get_pixmap(matrix=mat)
|
||||||
|
img_bytes = pix.tobytes("png")
|
||||||
|
doc.close()
|
||||||
|
return base64.b64encode(img_bytes).decode("utf-8")
|
||||||
|
|
||||||
|
except ImportError:
|
||||||
|
pass # Try pdf2image below
|
||||||
|
|
||||||
|
try:
|
||||||
|
from pdf2image import convert_from_bytes
|
||||||
|
|
||||||
|
images = convert_from_bytes(
|
||||||
|
pdf_bytes, dpi=dpi, first_page=page_num + 1, last_page=page_num + 1
|
||||||
|
)
|
||||||
|
if not images:
|
||||||
|
raise ValueError(f"pdf2image returned no images for page {page_num}")
|
||||||
|
buf = io.BytesIO()
|
||||||
|
images[0].save(buf, format="PNG")
|
||||||
|
return base64.b64encode(buf.getvalue()).decode("utf-8")
|
||||||
|
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(
|
||||||
|
"No PDF rendering library found. "
|
||||||
|
"Install PyMuPDF (pip install pymupdf) or pdf2image (pip install pdf2image)."
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# OCR request via sglang OpenAI-compatible API
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def run_ocr_request(
|
||||||
|
session: aiohttp.ClientSession,
|
||||||
|
api_url: str,
|
||||||
|
model: str,
|
||||||
|
image_b64: str,
|
||||||
|
text_prompt: str,
|
||||||
|
timeout: int,
|
||||||
|
) -> Tuple[str, float]:
|
||||||
|
"""
|
||||||
|
Send an image + text prompt to the sglang /v1/chat/completions endpoint.
|
||||||
|
|
||||||
|
Returns (ocr_text, latency_seconds).
|
||||||
|
On error returns ("ERROR: ...", -1.0).
|
||||||
|
"""
|
||||||
|
payload = {
|
||||||
|
"model": model,
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": f"data:image/png;base64,{image_b64}"},
|
||||||
|
},
|
||||||
|
{"type": "text", "text": text_prompt},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"max_tokens": 4096,
|
||||||
|
"temperature": 0.0,
|
||||||
|
}
|
||||||
|
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
try:
|
||||||
|
async with session.post(
|
||||||
|
api_url,
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||||
|
) as resp:
|
||||||
|
latency = time.perf_counter() - t0
|
||||||
|
if resp.status != 200:
|
||||||
|
body = await resp.text()
|
||||||
|
return f"ERROR: HTTP {resp.status} – {body[:200]}", -1.0
|
||||||
|
data = await resp.json()
|
||||||
|
text = data["choices"][0]["message"].get("content") or ""
|
||||||
|
return text, latency
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
return "ERROR: request timed out", -1.0
|
||||||
|
except Exception:
|
||||||
|
return f"ERROR: {traceback.format_exc(limit=3)}", -1.0
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Per-sample processing
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Local dataset loading (olmOCR-bench flat JSONL)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def load_jsonl_split(bench_dir: Path, split_name: str) -> List[dict]:
|
||||||
|
"""
|
||||||
|
Load test cases from a local olmOCR-bench JSONL file and group them by
|
||||||
|
(pdf, page) so each unique PDF page becomes one benchmark "sample".
|
||||||
|
|
||||||
|
Requires the dataset to have been downloaded first::
|
||||||
|
|
||||||
|
hf download --repo-type dataset \\
|
||||||
|
allenai/olmOCR-bench --local-dir ./olmOCR-bench
|
||||||
|
"""
|
||||||
|
jsonl_path = bench_dir / f"{split_name}.jsonl"
|
||||||
|
if not jsonl_path.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"JSONL not found: {jsonl_path}\n"
|
||||||
|
"Download the dataset first:\n"
|
||||||
|
" hf download --repo-type dataset "
|
||||||
|
"--resume-download allenai/olmOCR-bench --local-dir ./olmOCR-bench"
|
||||||
|
)
|
||||||
|
|
||||||
|
test_cases: List[dict] = []
|
||||||
|
with open(jsonl_path, encoding="utf-8") as fh:
|
||||||
|
for line in fh:
|
||||||
|
line = line.strip()
|
||||||
|
if line:
|
||||||
|
test_cases.append(json.loads(line))
|
||||||
|
|
||||||
|
# Group test cases by (pdf_relative_path, page)
|
||||||
|
pdf_dir = bench_dir / "pdfs"
|
||||||
|
groups: Dict[Tuple[str, int], dict] = {}
|
||||||
|
for tc in test_cases:
|
||||||
|
pdf_rel: str = tc["pdf"]
|
||||||
|
page: int = tc.get("page", 1)
|
||||||
|
key = (pdf_rel, page)
|
||||||
|
if key not in groups:
|
||||||
|
groups[key] = {
|
||||||
|
"pdf_path": str(pdf_dir / pdf_rel),
|
||||||
|
"pdf_rel": pdf_rel,
|
||||||
|
"page": page,
|
||||||
|
"tests": [],
|
||||||
|
}
|
||||||
|
groups[key]["tests"].append(tc)
|
||||||
|
|
||||||
|
return list(groups.values())
|
||||||
|
|
||||||
|
|
||||||
|
def _log_sample_error(label: str, error: str, debug: bool) -> None:
|
||||||
|
"""Print a sample error. Always shows a one-liner; full detail only in debug mode."""
|
||||||
|
short = error.splitlines()[0] if error else "unknown error"
|
||||||
|
print(f" [ERROR] {label}: {short}", file=sys.stderr)
|
||||||
|
if debug and len(error.splitlines()) > 1:
|
||||||
|
print(error, file=sys.stderr)
|
||||||
|
|
||||||
|
|
||||||
|
async def process_sample(
|
||||||
|
semaphore: asyncio.Semaphore,
|
||||||
|
session: aiohttp.ClientSession,
|
||||||
|
api_url: str,
|
||||||
|
args: BenchArgs,
|
||||||
|
text_prompt: str,
|
||||||
|
sample: dict,
|
||||||
|
) -> dict:
|
||||||
|
"""Process one benchmark sample (a single PDF page) and return its result dict."""
|
||||||
|
label = f"{sample.get('pdf_rel', '?')} page {sample.get('page', '?')}"
|
||||||
|
async with semaphore:
|
||||||
|
pdf_path = Path(sample["pdf_path"])
|
||||||
|
if not pdf_path.exists():
|
||||||
|
err = f"PDF not found: {pdf_path}"
|
||||||
|
_log_sample_error(label, err, args.debug)
|
||||||
|
return {"error": err, "test_results": [], "passed": 0, "total": 0}
|
||||||
|
|
||||||
|
try:
|
||||||
|
pdf_bytes = pdf_path.read_bytes()
|
||||||
|
except Exception as exc:
|
||||||
|
err = f"PDF read error: {exc}"
|
||||||
|
_log_sample_error(label, err, args.debug)
|
||||||
|
return {"error": err, "test_results": [], "passed": 0, "total": 0}
|
||||||
|
|
||||||
|
page_num = sample["page"] - 1 # convert 1-indexed → 0-indexed
|
||||||
|
|
||||||
|
try:
|
||||||
|
image_b64 = pdf_page_to_base64_png(
|
||||||
|
pdf_bytes, page_num=page_num, dpi=args.render_dpi
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
err = f"PDF render error: {exc}"
|
||||||
|
if args.debug:
|
||||||
|
err += "\n" + traceback.format_exc()
|
||||||
|
_log_sample_error(label, err, args.debug)
|
||||||
|
return {"error": err, "test_results": [], "passed": 0, "total": 0}
|
||||||
|
|
||||||
|
ocr_text, latency = await run_ocr_request(
|
||||||
|
session, api_url, args.model, image_b64, text_prompt, args.request_timeout
|
||||||
|
)
|
||||||
|
|
||||||
|
if ocr_text.startswith("ERROR:"):
|
||||||
|
_log_sample_error(label, ocr_text, args.debug)
|
||||||
|
return {"error": ocr_text, "test_results": [], "passed": 0, "total": 0}
|
||||||
|
|
||||||
|
tests = sample["tests"]
|
||||||
|
test_results = evaluate_olmocr_tests(tests, ocr_text)
|
||||||
|
|
||||||
|
# Accuracy debug: show input expectations and full OCR output
|
||||||
|
if args.debug_accuracy:
|
||||||
|
sep = "-" * 72
|
||||||
|
print(f"\n{sep}", flush=True)
|
||||||
|
print(f"[INPUT ] PDF : {pdf_path}", flush=True)
|
||||||
|
print(
|
||||||
|
f"[INPUT ] Page : {sample['page']} | {len(tests)} test(s)", flush=True
|
||||||
|
)
|
||||||
|
for i, t in enumerate(tests):
|
||||||
|
ttype = t.get("type", "?")
|
||||||
|
if ttype in ("present", "absent", "text_presence", "text_absence"):
|
||||||
|
expected = t.get("text", "")
|
||||||
|
elif ttype in ("math", "math_formula_accuracy"):
|
||||||
|
expected = t.get("math") or t.get("latex", "")
|
||||||
|
elif ttype in ("order", "natural_reading_order"):
|
||||||
|
expected = (
|
||||||
|
f"before={t.get('before', '')!r} after={t.get('after', '')!r}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
expected = str(
|
||||||
|
{
|
||||||
|
k: v
|
||||||
|
for k, v in t.items()
|
||||||
|
if k not in ("pdf", "page", "id", "type", "url", "checked")
|
||||||
|
}
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[INPUT ] [{i+1}] type={ttype!r:12s} expected: {expected[:120]}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
# Print OCR output (truncate long outputs)
|
||||||
|
ocr_preview = (
|
||||||
|
ocr_text
|
||||||
|
if len(ocr_text) <= 800
|
||||||
|
else ocr_text[:800] + f"\n... [{len(ocr_text)} chars total, truncated]"
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[OUTPUT] OCR text ({len(ocr_text)} chars, latency={latency:.2f}s):",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(ocr_preview, flush=True)
|
||||||
|
|
||||||
|
# Log individual test failures in debug mode
|
||||||
|
if args.debug or args.debug_accuracy:
|
||||||
|
for tr in test_results:
|
||||||
|
status = "PASS" if tr.get("passed") else "FAIL"
|
||||||
|
detail = tr.get("error", "") if not tr.get("passed") else ""
|
||||||
|
suffix = f" | {detail}" if detail else ""
|
||||||
|
print(
|
||||||
|
f"[RESULT] [{status}] type={tr.get('type')!r}{suffix}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
if args.debug_accuracy:
|
||||||
|
print(sep, flush=True)
|
||||||
|
|
||||||
|
result: dict = {
|
||||||
|
"pdf": sample["pdf_rel"],
|
||||||
|
"page": sample["page"],
|
||||||
|
"latency": round(latency, 3),
|
||||||
|
"test_results": test_results,
|
||||||
|
"passed": sum(1 for r in test_results if r.get("passed")),
|
||||||
|
"total": len(test_results),
|
||||||
|
# Store expected values so the HTML report can render them
|
||||||
|
"test_inputs": [
|
||||||
|
{
|
||||||
|
"type": t.get("type"),
|
||||||
|
"math": t.get("math") or t.get("latex", ""),
|
||||||
|
"text": t.get("text", ""),
|
||||||
|
"before": t.get("before", ""),
|
||||||
|
"after": t.get("after", ""),
|
||||||
|
# table-specific fields
|
||||||
|
"cell": t.get("cell", ""),
|
||||||
|
"up": t.get("up"),
|
||||||
|
"down": t.get("down"),
|
||||||
|
"left": t.get("left"),
|
||||||
|
"right": t.get("right"),
|
||||||
|
"top_heading": t.get("top_heading"),
|
||||||
|
"left_heading": t.get("left_heading"),
|
||||||
|
}
|
||||||
|
for t in tests
|
||||||
|
],
|
||||||
|
}
|
||||||
|
if args.save_raw_outputs:
|
||||||
|
result["ocr_output"] = ocr_text
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Split-level runner
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def run_split(
|
||||||
|
split_name: str,
|
||||||
|
dataset,
|
||||||
|
args: BenchArgs,
|
||||||
|
api_url: str,
|
||||||
|
text_prompt: str,
|
||||||
|
) -> dict:
|
||||||
|
"""Evaluate one olmOCR-bench split; return aggregated results dict."""
|
||||||
|
samples = list(dataset)
|
||||||
|
if args.max_samples > 0:
|
||||||
|
samples = samples[: args.max_samples]
|
||||||
|
|
||||||
|
semaphore = asyncio.Semaphore(args.concurrency)
|
||||||
|
sample_results: List[dict] = []
|
||||||
|
_first_error: List[str] = [] # capture first error for summary
|
||||||
|
|
||||||
|
connector = aiohttp.TCPConnector(limit=args.concurrency + 4)
|
||||||
|
async with aiohttp.ClientSession(connector=connector) as session:
|
||||||
|
tasks = [
|
||||||
|
process_sample(semaphore, session, api_url, args, text_prompt, sample)
|
||||||
|
for sample in samples
|
||||||
|
]
|
||||||
|
for future in atqdm(
|
||||||
|
asyncio.as_completed(tasks),
|
||||||
|
total=len(tasks),
|
||||||
|
desc=f" [{split_name}]",
|
||||||
|
leave=True,
|
||||||
|
):
|
||||||
|
result = await future
|
||||||
|
if "error" in result and not _first_error:
|
||||||
|
_first_error.append(result["error"])
|
||||||
|
sample_results.append(result)
|
||||||
|
|
||||||
|
agg = aggregate_results(split_name, sample_results)
|
||||||
|
agg["samples"] = sample_results # include per-sample data for report generation
|
||||||
|
|
||||||
|
# Always surface error summary so silent 0/0 can't happen
|
||||||
|
error_count = agg.get("error_samples", 0)
|
||||||
|
if error_count > 0:
|
||||||
|
first = _first_error[0] if _first_error else "(unknown)"
|
||||||
|
short = first.splitlines()[0]
|
||||||
|
print(
|
||||||
|
f" WARNING: {error_count}/{len(samples)} samples errored and were skipped.",
|
||||||
|
file=sys.stderr,
|
||||||
|
)
|
||||||
|
print(f" First error: {short}", file=sys.stderr)
|
||||||
|
if not args.debug:
|
||||||
|
print(
|
||||||
|
" Re-run with --debug for full per-sample error details.",
|
||||||
|
file=sys.stderr,
|
||||||
|
)
|
||||||
|
|
||||||
|
return agg
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Entry point
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> BenchArgs:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Benchmark OCR VLMs on olmOCR-bench via sglang",
|
||||||
|
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||||
|
)
|
||||||
|
BenchArgs.add_cli_args(parser)
|
||||||
|
ns = parser.parse_args()
|
||||||
|
return BenchArgs.from_cli_args(ns)
|
||||||
|
|
||||||
|
|
||||||
|
async def main() -> None:
|
||||||
|
args = parse_args()
|
||||||
|
|
||||||
|
api_url = f"http://{args.host}:{args.port}/v1/chat/completions"
|
||||||
|
text_prompt = (
|
||||||
|
_PROMPT_MARKDOWN if args.prompt_mode == "markdown" else _PROMPT_FREE_OCR
|
||||||
|
)
|
||||||
|
splits_to_run = OLMOCR_BENCH_SPLITS if args.split == "all" else [args.split]
|
||||||
|
|
||||||
|
print("=" * 60)
|
||||||
|
print(f" OCR Accuracy Benchmark – olmOCR-bench")
|
||||||
|
print("=" * 60)
|
||||||
|
print(f" Model : {args.model}")
|
||||||
|
print(f" Server : {api_url}")
|
||||||
|
print(f" Prompt mode : {args.prompt_mode}")
|
||||||
|
print(f" Splits : {', '.join(splits_to_run)}")
|
||||||
|
print(f" Concurrency : {args.concurrency}")
|
||||||
|
print(f" Bench dir : {args.bench_dir}")
|
||||||
|
print(f" Output dir : {args.output_dir}")
|
||||||
|
if args.debug:
|
||||||
|
print(f" Debug mode : ON (errors)")
|
||||||
|
if args.debug_accuracy:
|
||||||
|
print(f" Debug mode : ON (accuracy — input/output per sample)")
|
||||||
|
print("=" * 60)
|
||||||
|
|
||||||
|
await preflight_check(api_url, args.model, args.debug)
|
||||||
|
|
||||||
|
bench_dir = Path(args.bench_dir)
|
||||||
|
if not bench_dir.exists():
|
||||||
|
raise SystemExit(
|
||||||
|
f"Benchmark directory not found: {bench_dir}\n"
|
||||||
|
"Download the dataset first:\n"
|
||||||
|
" hf download --repo-type dataset "
|
||||||
|
"allenai/olmOCR-bench --local-dir ./olmOCR-bench"
|
||||||
|
)
|
||||||
|
|
||||||
|
os.makedirs(args.output_dir, exist_ok=True)
|
||||||
|
all_results: Dict[str, dict] = {}
|
||||||
|
|
||||||
|
for split in splits_to_run:
|
||||||
|
print(f"\nLoading split '{split}' from {bench_dir} …")
|
||||||
|
try:
|
||||||
|
samples = load_jsonl_split(bench_dir, split)
|
||||||
|
except FileNotFoundError as exc:
|
||||||
|
print(f" WARNING: {exc}")
|
||||||
|
continue
|
||||||
|
except Exception as exc:
|
||||||
|
print(f" WARNING: could not load split '{split}': {exc}")
|
||||||
|
continue
|
||||||
|
|
||||||
|
n = (
|
||||||
|
len(samples)
|
||||||
|
if args.max_samples <= 0
|
||||||
|
else min(len(samples), args.max_samples)
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f" {n} PDF pages to evaluate ({sum(len(s['tests']) for s in samples[:n])} tests) …"
|
||||||
|
)
|
||||||
|
|
||||||
|
split_result = await run_split(split, samples, args, api_url, text_prompt)
|
||||||
|
all_results[split] = split_result
|
||||||
|
|
||||||
|
# Save per-split JSON
|
||||||
|
out_path = os.path.join(args.output_dir, f"{split}.json")
|
||||||
|
with open(out_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(split_result, f, indent=2, ensure_ascii=False)
|
||||||
|
print(
|
||||||
|
f" Score: {split_result['overall_score']:.1f}% "
|
||||||
|
f"({split_result['total_passed']}/{split_result['total_tests']} tests passed)"
|
||||||
|
)
|
||||||
|
print(f" Saved → {out_path}")
|
||||||
|
|
||||||
|
if not all_results:
|
||||||
|
print("No results collected – exiting.")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Print final table
|
||||||
|
print_results_table(all_results)
|
||||||
|
|
||||||
|
# Save summary
|
||||||
|
summary_path = os.path.join(args.output_dir, "summary.json")
|
||||||
|
with open(summary_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(all_results, f, indent=2, ensure_ascii=False)
|
||||||
|
print(f"\nFull summary saved → {summary_path}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
@@ -0,0 +1,631 @@
|
|||||||
|
"""
|
||||||
|
Evaluation utilities for the OCR benchmark (olmOCR-bench test classes).
|
||||||
|
|
||||||
|
Implements:
|
||||||
|
- text_presence : short text segment must be present in OCR output
|
||||||
|
- text_absence : text (headers/footers/page numbers) must NOT appear
|
||||||
|
- natural_reading_order : two text spans must appear in correct relative order
|
||||||
|
- table_accuracy : cell values with correct neighbor relationships (Markdown + HTML)
|
||||||
|
- math_formula_accuracy : LaTeX key-symbol token matching (simplified; no KaTeX/playwright)
|
||||||
|
|
||||||
|
Also provides:
|
||||||
|
- normalized_edit_distance() for OmniDocBench-style text quality measurement
|
||||||
|
- aggregate_results() / print_results_table() for summary reporting
|
||||||
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
from difflib import SequenceMatcher
|
||||||
|
from html.parser import HTMLParser
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
# ── Unicode normalization ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
_HYPHEN_RE = re.compile(
|
||||||
|
r"[\u2010\u2011\u2012\u2013\u2014\u2015\u2212\uFE58\uFE63\uFF0D]"
|
||||||
|
)
|
||||||
|
_DQUOTE_RE = re.compile(
|
||||||
|
r"[\u00AB\u00BB\u201C\u201D\u201E\u201F\u2033\u2036\u276E\u276F\u3003\uFF02]"
|
||||||
|
)
|
||||||
|
_SQUOTE_RE = re.compile(
|
||||||
|
r"[\u2018\u2019\u201A\u201B\u2032\u2035\u2039\u203A\u2C8D\uFF07]"
|
||||||
|
)
|
||||||
|
_MARKDOWN_RE = re.compile(r"(\*{1,3}|_{1,3}|`{1,3}|~~|#{1,6}\s?)")
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
"""Apply olmOCR-bench standard Unicode normalization."""
|
||||||
|
text = unicodedata.normalize("NFC", text)
|
||||||
|
text = _HYPHEN_RE.sub("-", text)
|
||||||
|
text = _DQUOTE_RE.sub('"', text)
|
||||||
|
text = _SQUOTE_RE.sub("'", text)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def strip_markdown(text: str) -> str:
|
||||||
|
"""Remove Markdown syntax markers for soft matching."""
|
||||||
|
return _MARKDOWN_RE.sub("", text)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Matching helpers ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def fuzzy_contains(needle: str, haystack: str, threshold: float = 0.85) -> bool:
|
||||||
|
"""Check if needle appears in haystack using fuzzy sliding-window matching."""
|
||||||
|
needle = normalize_text(strip_markdown(needle).strip())
|
||||||
|
haystack = normalize_text(strip_markdown(haystack))
|
||||||
|
|
||||||
|
# Fast exact check first
|
||||||
|
if needle.lower() in haystack.lower():
|
||||||
|
return True
|
||||||
|
|
||||||
|
n = len(needle)
|
||||||
|
if n == 0:
|
||||||
|
return True
|
||||||
|
|
||||||
|
step = max(1, n // 4)
|
||||||
|
for i in range(0, max(1, len(haystack) - n + 1), step):
|
||||||
|
window = haystack[i : i + n]
|
||||||
|
ratio = SequenceMatcher(None, needle.lower(), window.lower()).ratio()
|
||||||
|
if ratio >= threshold:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def exact_contains(needle: str, haystack: str, case_sensitive: bool = True) -> bool:
|
||||||
|
"""Check if needle appears exactly in haystack (after normalization)."""
|
||||||
|
needle = normalize_text(strip_markdown(needle).strip())
|
||||||
|
haystack = normalize_text(strip_markdown(haystack))
|
||||||
|
if not case_sensitive:
|
||||||
|
return needle.lower() in haystack.lower()
|
||||||
|
return needle in haystack
|
||||||
|
|
||||||
|
|
||||||
|
def _get_words_slice(text: str, first_n: Optional[int], last_n: Optional[int]) -> str:
|
||||||
|
"""Return the first or last N whitespace-separated words of text."""
|
||||||
|
if first_n is None and last_n is None:
|
||||||
|
return text
|
||||||
|
words = text.split()
|
||||||
|
if first_n is not None:
|
||||||
|
return " ".join(words[:first_n])
|
||||||
|
if last_n is not None:
|
||||||
|
return " ".join(words[-last_n:])
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
# ── olmOCR-bench test evaluators ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def eval_text_presence(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""Evaluate a present/text_presence test: target text must appear in OCR output.
|
||||||
|
|
||||||
|
Supports olmOCR-bench flat schema (max_diffs, first_n, last_n) and the
|
||||||
|
legacy nested-position schema.
|
||||||
|
"""
|
||||||
|
needle = test.get("text", "")
|
||||||
|
max_diffs = test.get("max_diffs", 0)
|
||||||
|
case_sensitive = test.get("case_sensitive", True)
|
||||||
|
first_n = test.get("first_n", None)
|
||||||
|
last_n = test.get("last_n", None)
|
||||||
|
|
||||||
|
haystack = _get_words_slice(ocr_output, first_n, last_n)
|
||||||
|
|
||||||
|
if max_diffs == 0:
|
||||||
|
return exact_contains(needle, haystack, case_sensitive=case_sensitive)
|
||||||
|
# Fuzzy: compute similarity threshold from allowed diffs
|
||||||
|
n = max(1, len(needle))
|
||||||
|
threshold = max(0.6, (n - max_diffs) / n)
|
||||||
|
return fuzzy_contains(needle, haystack, threshold=threshold)
|
||||||
|
|
||||||
|
|
||||||
|
def eval_text_absence(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""Evaluate an absent/text_absence test: target text must NOT appear in OCR output.
|
||||||
|
|
||||||
|
Supports olmOCR-bench flat schema (max_diffs, first_n, last_n) and the
|
||||||
|
legacy nested-position schema.
|
||||||
|
"""
|
||||||
|
needle = test.get("text", "")
|
||||||
|
max_diffs = test.get("max_diffs", 0)
|
||||||
|
case_sensitive = test.get("case_sensitive", False)
|
||||||
|
first_n = test.get("first_n", None)
|
||||||
|
last_n = test.get("last_n", None)
|
||||||
|
|
||||||
|
haystack = _get_words_slice(ocr_output, first_n, last_n)
|
||||||
|
|
||||||
|
if max_diffs == 0:
|
||||||
|
present = exact_contains(needle, haystack, case_sensitive=case_sensitive)
|
||||||
|
else:
|
||||||
|
n = max(1, len(needle))
|
||||||
|
threshold = max(0.6, (n - max_diffs) / n)
|
||||||
|
present = fuzzy_contains(needle, haystack, threshold=threshold)
|
||||||
|
return not present
|
||||||
|
|
||||||
|
|
||||||
|
def eval_reading_order(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""Evaluate an order/natural_reading_order test: 'before' text must precede 'after'.
|
||||||
|
|
||||||
|
Uses max_diffs to choose exact vs fuzzy matching.
|
||||||
|
"""
|
||||||
|
before_text = test.get("before", "")
|
||||||
|
after_text = test.get("after", "")
|
||||||
|
max_diffs = test.get("max_diffs", 0)
|
||||||
|
fuzzy = max_diffs > 0
|
||||||
|
|
||||||
|
output_norm = normalize_text(strip_markdown(ocr_output))
|
||||||
|
|
||||||
|
def find_approx_pos(needle: str, text: str) -> int:
|
||||||
|
needle = normalize_text(strip_markdown(needle).strip())
|
||||||
|
n = len(needle)
|
||||||
|
if n == 0:
|
||||||
|
return 0
|
||||||
|
# Exact first
|
||||||
|
idx = text.lower().find(needle.lower())
|
||||||
|
if idx != -1:
|
||||||
|
return idx
|
||||||
|
# Fuzzy fallback
|
||||||
|
step = max(1, n // 4)
|
||||||
|
best_pos, best_ratio = -1, 0.0
|
||||||
|
for i in range(0, max(1, len(text) - n + 1), step):
|
||||||
|
window = text[i : i + n]
|
||||||
|
ratio = SequenceMatcher(None, needle.lower(), window.lower()).ratio()
|
||||||
|
if ratio > best_ratio:
|
||||||
|
best_ratio = ratio
|
||||||
|
best_pos = i
|
||||||
|
return best_pos if best_ratio >= 0.80 else -1
|
||||||
|
|
||||||
|
if fuzzy:
|
||||||
|
pos_before = find_approx_pos(before_text, output_norm)
|
||||||
|
pos_after = find_approx_pos(after_text, output_norm)
|
||||||
|
else:
|
||||||
|
b = normalize_text(strip_markdown(before_text).strip())
|
||||||
|
a = normalize_text(strip_markdown(after_text).strip())
|
||||||
|
pos_before = output_norm.find(b)
|
||||||
|
pos_after = output_norm.find(a)
|
||||||
|
|
||||||
|
if pos_before == -1 or pos_after == -1:
|
||||||
|
return False
|
||||||
|
return pos_before < pos_after
|
||||||
|
|
||||||
|
|
||||||
|
# ── Table parsing helpers ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_markdown_table(text: str) -> List[List[str]]:
|
||||||
|
"""Parse a Markdown table into a list-of-rows, each row a list of cells."""
|
||||||
|
rows: List[List[str]] = []
|
||||||
|
for line in text.splitlines():
|
||||||
|
stripped = line.strip()
|
||||||
|
if "|" not in stripped:
|
||||||
|
continue
|
||||||
|
# Skip separator rows like |---|---|
|
||||||
|
if re.match(r"^\|?[-:| ]+\|?$", stripped):
|
||||||
|
continue
|
||||||
|
cells = [c.strip() for c in stripped.strip("|").split("|")]
|
||||||
|
if cells:
|
||||||
|
rows.append(cells)
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
class _HTMLTableParser(HTMLParser):
|
||||||
|
"""Minimal HTML table parser (does not handle colspan/rowspan)."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.rows: List[List[str]] = []
|
||||||
|
self._current_row: List[str] = []
|
||||||
|
self._current_cell: str = ""
|
||||||
|
self._in_cell: bool = False
|
||||||
|
|
||||||
|
def handle_starttag(self, tag: str, attrs) -> None:
|
||||||
|
if tag == "tr":
|
||||||
|
self._current_row = []
|
||||||
|
elif tag in ("td", "th"):
|
||||||
|
self._in_cell = True
|
||||||
|
self._current_cell = ""
|
||||||
|
|
||||||
|
def handle_endtag(self, tag: str) -> None:
|
||||||
|
if tag in ("td", "th"):
|
||||||
|
self._current_row.append(self._current_cell.strip())
|
||||||
|
self._in_cell = False
|
||||||
|
elif tag == "tr" and self._current_row:
|
||||||
|
self.rows.append(self._current_row)
|
||||||
|
|
||||||
|
def handle_data(self, data: str) -> None:
|
||||||
|
if self._in_cell:
|
||||||
|
self._current_cell += data
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_tables(ocr_output: str) -> List[List[List[str]]]:
|
||||||
|
"""Extract all tables (HTML + Markdown) from OCR output."""
|
||||||
|
tables: List[List[List[str]]] = []
|
||||||
|
|
||||||
|
# HTML tables
|
||||||
|
for match in re.finditer(
|
||||||
|
r"<table[^>]*>.*?</table>", ocr_output, re.DOTALL | re.IGNORECASE
|
||||||
|
):
|
||||||
|
parser = _HTMLTableParser()
|
||||||
|
parser.feed(match.group(0))
|
||||||
|
if parser.rows:
|
||||||
|
tables.append(parser.rows)
|
||||||
|
|
||||||
|
# Markdown tables
|
||||||
|
md_pattern = re.compile(
|
||||||
|
r"(\|[^\n]+\|\n(?:\|[-:| ]+\|\n)?(?:\|[^\n]+\|?\n?)+)", re.MULTILINE
|
||||||
|
)
|
||||||
|
for match in md_pattern.finditer(ocr_output):
|
||||||
|
rows = _parse_markdown_table(match.group(0))
|
||||||
|
if len(rows) >= 2: # At least header + one data row
|
||||||
|
tables.append(rows)
|
||||||
|
|
||||||
|
return tables
|
||||||
|
|
||||||
|
|
||||||
|
def eval_table_flat(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""
|
||||||
|
Evaluate a flat-schema olmOCR-bench 'table' test.
|
||||||
|
|
||||||
|
Schema fields:
|
||||||
|
cell – text of the target cell to locate
|
||||||
|
up/down/left/right – expected text of the neighbor in that direction (null = skip)
|
||||||
|
top_heading – expected column heading (row 0, same column)
|
||||||
|
left_heading – expected row heading (column 0, same row)
|
||||||
|
|
||||||
|
All non-null fields must fuzzy-match for the test to pass.
|
||||||
|
|
||||||
|
Supports:
|
||||||
|
1. Structured tables (HTML with <tr>/<td>, Markdown) – full positional check
|
||||||
|
2. Flat <table>content</table> blocks (DeepSeek-OCR-2 format) – text presence
|
||||||
|
fallback when no rows are parseable from HTML
|
||||||
|
"""
|
||||||
|
cell_text = test.get("cell", "")
|
||||||
|
directional = {
|
||||||
|
"up": test.get("up"),
|
||||||
|
"down": test.get("down"),
|
||||||
|
"left": test.get("left"),
|
||||||
|
"right": test.get("right"),
|
||||||
|
}
|
||||||
|
top_heading = test.get("top_heading")
|
||||||
|
left_heading = test.get("left_heading")
|
||||||
|
checks = [(d, v) for d, v in directional.items() if v is not None]
|
||||||
|
|
||||||
|
# ── 1. Structured tables (HTML <tr>/<td> or Markdown) ────────────────────
|
||||||
|
for rows in _extract_tables(ocr_output):
|
||||||
|
if not rows:
|
||||||
|
continue
|
||||||
|
header_row = rows[0]
|
||||||
|
for r_idx, row in enumerate(rows):
|
||||||
|
for c_idx, cell in enumerate(row):
|
||||||
|
if not fuzzy_contains(cell_text, cell, threshold=0.85):
|
||||||
|
continue
|
||||||
|
# Verify all directional neighbors
|
||||||
|
all_ok = True
|
||||||
|
for direction, expected in checks:
|
||||||
|
if direction == "up" and r_idx > 0:
|
||||||
|
prev_row = rows[r_idx - 1]
|
||||||
|
nb = prev_row[c_idx] if c_idx < len(prev_row) else ""
|
||||||
|
if not fuzzy_contains(expected, nb, threshold=0.85):
|
||||||
|
all_ok = False
|
||||||
|
break
|
||||||
|
elif direction == "down" and r_idx < len(rows) - 1:
|
||||||
|
next_row = rows[r_idx + 1]
|
||||||
|
nb = next_row[c_idx] if c_idx < len(next_row) else ""
|
||||||
|
if not fuzzy_contains(expected, nb, threshold=0.85):
|
||||||
|
all_ok = False
|
||||||
|
break
|
||||||
|
elif direction == "left" and c_idx > 0:
|
||||||
|
nb = row[c_idx - 1]
|
||||||
|
if not fuzzy_contains(expected, nb, threshold=0.85):
|
||||||
|
all_ok = False
|
||||||
|
break
|
||||||
|
elif direction == "right" and c_idx < len(row) - 1:
|
||||||
|
nb = row[c_idx + 1]
|
||||||
|
if not fuzzy_contains(expected, nb, threshold=0.85):
|
||||||
|
all_ok = False
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
all_ok = False
|
||||||
|
break # expected neighbor out of bounds
|
||||||
|
if not all_ok:
|
||||||
|
continue
|
||||||
|
# Verify top_heading (column header, row 0)
|
||||||
|
if top_heading is not None:
|
||||||
|
th = header_row[c_idx] if c_idx < len(header_row) else ""
|
||||||
|
if not fuzzy_contains(top_heading, th, threshold=0.85):
|
||||||
|
continue
|
||||||
|
# Verify left_heading (first cell of same row)
|
||||||
|
if left_heading is not None:
|
||||||
|
lh = row[0] if row else ""
|
||||||
|
if not fuzzy_contains(left_heading, lh, threshold=0.85):
|
||||||
|
continue
|
||||||
|
return True
|
||||||
|
|
||||||
|
# ── 2. Flat <table>…</table> fallback (DeepSeek-OCR-2 format) ────────────
|
||||||
|
# The model emits <table>AllCellsConcatenated</table> without <tr>/<td>.
|
||||||
|
# Fall back to checking that cell + all non-null headings/neighbors appear
|
||||||
|
# somewhere within the same table block.
|
||||||
|
flat_blocks = re.findall(
|
||||||
|
r"<table[^>]*>(.*?)</table>", ocr_output, re.DOTALL | re.IGNORECASE
|
||||||
|
)
|
||||||
|
for flat_text in flat_blocks:
|
||||||
|
# Strip any residual HTML tags (e.g. inline <br>) and bounding-box annotations
|
||||||
|
flat_clean = re.sub(r"<[^>]+>", " ", flat_text)
|
||||||
|
flat_clean = re.sub(r"\[\[\d+,\s*\d+,\s*\d+,\s*\d+\]\]", " ", flat_clean)
|
||||||
|
if not fuzzy_contains(cell_text, flat_clean, threshold=0.85):
|
||||||
|
continue
|
||||||
|
if top_heading is not None and not fuzzy_contains(
|
||||||
|
top_heading, flat_clean, threshold=0.85
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if left_heading is not None and not fuzzy_contains(
|
||||||
|
left_heading, flat_clean, threshold=0.85
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
all_ok = all(
|
||||||
|
fuzzy_contains(expected, flat_clean, threshold=0.85)
|
||||||
|
for _, expected in checks
|
||||||
|
)
|
||||||
|
if all_ok:
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def eval_baseline(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""
|
||||||
|
Evaluate a baseline/sanity test.
|
||||||
|
|
||||||
|
When check_disallowed_characters is False (the common case), always passes.
|
||||||
|
When True, checks that the OCR output contains no non-printable control
|
||||||
|
characters (excluding normal whitespace).
|
||||||
|
"""
|
||||||
|
if not test.get("check_disallowed_characters", False):
|
||||||
|
return True
|
||||||
|
for ch in ocr_output:
|
||||||
|
cat = unicodedata.category(ch)
|
||||||
|
if cat.startswith("C") and ch not in ("\n", "\t", "\r", " "):
|
||||||
|
return False
|
||||||
|
return bool(ocr_output.strip()) # also fail if completely empty
|
||||||
|
|
||||||
|
|
||||||
|
def eval_table_accuracy(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""
|
||||||
|
Evaluate a table_accuracy test (legacy nested schema).
|
||||||
|
|
||||||
|
Checks that a cell with ``cell_text`` exists in a table and that its
|
||||||
|
neighbor in the specified ``relationship`` (above/below/left/right)
|
||||||
|
contains ``neighbor_text``.
|
||||||
|
"""
|
||||||
|
cell_text = test.get("cell_text", "")
|
||||||
|
neighbor_text = test.get("neighbor_text", "")
|
||||||
|
relationship = test.get("relationship", "")
|
||||||
|
|
||||||
|
for rows in _extract_tables(ocr_output):
|
||||||
|
for r_idx, row in enumerate(rows):
|
||||||
|
for c_idx, cell in enumerate(row):
|
||||||
|
if not fuzzy_contains(cell_text, cell, threshold=0.88):
|
||||||
|
continue
|
||||||
|
# Found the target cell — check its neighbor
|
||||||
|
if relationship == "above" and r_idx > 0:
|
||||||
|
prev_row = rows[r_idx - 1]
|
||||||
|
nb = prev_row[c_idx] if c_idx < len(prev_row) else ""
|
||||||
|
if fuzzy_contains(neighbor_text, nb, threshold=0.88):
|
||||||
|
return True
|
||||||
|
elif relationship == "below" and r_idx < len(rows) - 1:
|
||||||
|
next_row = rows[r_idx + 1]
|
||||||
|
nb = next_row[c_idx] if c_idx < len(next_row) else ""
|
||||||
|
if fuzzy_contains(neighbor_text, nb, threshold=0.88):
|
||||||
|
return True
|
||||||
|
elif relationship == "left" and c_idx > 0:
|
||||||
|
nb = row[c_idx - 1]
|
||||||
|
if fuzzy_contains(neighbor_text, nb, threshold=0.88):
|
||||||
|
return True
|
||||||
|
elif relationship == "right" and c_idx < len(row) - 1:
|
||||||
|
nb = row[c_idx + 1]
|
||||||
|
if fuzzy_contains(neighbor_text, nb, threshold=0.88):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# Math formula evaluation ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
_MATH_REGION_RE = re.compile(
|
||||||
|
r"\$\$[\s\S]*?\$\$" # $$ block $$
|
||||||
|
r"|\$[^$\n]+?\$" # inline $...$
|
||||||
|
r"|\\?\\\[[\s\S]*?\\?\\\]" # \[...\]
|
||||||
|
r"|\\?\\\([\s\S]*?\\?\\\)", # \(...\)
|
||||||
|
)
|
||||||
|
_LATEX_TOKEN_RE = re.compile(r"\\[a-zA-Z]+|[a-zA-Z0-9]|[+\-*/=<>^_{}()\[\]]")
|
||||||
|
|
||||||
|
|
||||||
|
def eval_math_formula_accuracy(test: dict, ocr_output: str) -> bool:
|
||||||
|
"""
|
||||||
|
Simplified math formula accuracy check.
|
||||||
|
|
||||||
|
Checks that the key symbol tokens from a LaTeX expression appear in
|
||||||
|
math-delimited regions of the OCR output.
|
||||||
|
|
||||||
|
Note: Full KaTeX bounding-box matching (as used by the official
|
||||||
|
olmOCR-bench) requires playwright and is not performed here.
|
||||||
|
"""
|
||||||
|
latex = (test.get("math") or test.get("latex") or "").strip()
|
||||||
|
if not latex:
|
||||||
|
return False
|
||||||
|
|
||||||
|
math_text = " ".join(m.group(0) for m in _MATH_REGION_RE.finditer(ocr_output))
|
||||||
|
if not math_text:
|
||||||
|
# Fall back to full output if no delimited regions found
|
||||||
|
math_text = ocr_output
|
||||||
|
|
||||||
|
tokens = _LATEX_TOKEN_RE.findall(latex)
|
||||||
|
if not tokens:
|
||||||
|
return False
|
||||||
|
|
||||||
|
present = sum(1 for t in tokens if t in math_text)
|
||||||
|
return present / len(tokens) >= 0.70
|
||||||
|
|
||||||
|
|
||||||
|
# ── Main dispatcher ───────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
_TEST_EVALUATORS = {
|
||||||
|
# olmOCR-bench flat-JSONL type names
|
||||||
|
"present": eval_text_presence,
|
||||||
|
"absent": eval_text_absence,
|
||||||
|
"order": eval_reading_order,
|
||||||
|
"math": eval_math_formula_accuracy,
|
||||||
|
"table": eval_table_flat,
|
||||||
|
"baseline": eval_baseline,
|
||||||
|
# Legacy / aliased names
|
||||||
|
"text_presence": eval_text_presence,
|
||||||
|
"text_absence": eval_text_absence,
|
||||||
|
"natural_reading_order": eval_reading_order,
|
||||||
|
"table_accuracy": eval_table_accuracy,
|
||||||
|
"math_formula_accuracy": eval_math_formula_accuracy,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate_olmocr_tests(tests: List[dict], ocr_output: str) -> List[dict]:
|
||||||
|
"""Run all olmOCR-bench unit tests against OCR output; return per-test results."""
|
||||||
|
results: List[dict] = []
|
||||||
|
for test in tests:
|
||||||
|
test_type = test.get("type", "")
|
||||||
|
evaluator = _TEST_EVALUATORS.get(test_type)
|
||||||
|
if evaluator is None:
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"type": test_type,
|
||||||
|
"passed": False,
|
||||||
|
"error": f"Unknown type: {test_type}",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
passed = bool(evaluator(test, ocr_output))
|
||||||
|
except Exception as exc:
|
||||||
|
results.append({"type": test_type, "passed": False, "error": str(exc)})
|
||||||
|
continue
|
||||||
|
results.append({"type": test_type, "passed": passed})
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
# ── Aggregation & reporting ───────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def aggregate_results(split_name: str, sample_results: List[dict]) -> dict:
|
||||||
|
"""Aggregate per-sample results into split-level statistics."""
|
||||||
|
by_type: Dict[str, Dict[str, int]] = {}
|
||||||
|
total_passed = 0
|
||||||
|
total_tests = 0
|
||||||
|
error_count = 0
|
||||||
|
|
||||||
|
for sample in sample_results:
|
||||||
|
if "error" in sample and not sample.get("test_results"):
|
||||||
|
error_count += 1
|
||||||
|
continue
|
||||||
|
for tr in sample.get("test_results", []):
|
||||||
|
t = tr.get("type", "unknown")
|
||||||
|
by_type.setdefault(t, {"passed": 0, "total": 0})
|
||||||
|
by_type[t]["total"] += 1
|
||||||
|
total_tests += 1
|
||||||
|
if tr.get("passed"):
|
||||||
|
by_type[t]["passed"] += 1
|
||||||
|
total_passed += 1
|
||||||
|
|
||||||
|
type_scores = {
|
||||||
|
t: round(100.0 * v["passed"] / v["total"], 1) if v["total"] > 0 else 0.0
|
||||||
|
for t, v in by_type.items()
|
||||||
|
}
|
||||||
|
overall = round(100.0 * total_passed / total_tests, 1) if total_tests > 0 else 0.0
|
||||||
|
|
||||||
|
return {
|
||||||
|
"split": split_name,
|
||||||
|
"total_samples": len(sample_results),
|
||||||
|
"error_samples": error_count,
|
||||||
|
"total_tests": total_tests,
|
||||||
|
"total_passed": total_passed,
|
||||||
|
"overall_score": overall,
|
||||||
|
"by_type": type_scores,
|
||||||
|
"by_type_counts": by_type,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def print_results_table(all_results: Dict[str, dict]) -> None:
|
||||||
|
"""Print a formatted results summary table to stdout."""
|
||||||
|
sep = "=" * 70
|
||||||
|
print(f"\n{sep}")
|
||||||
|
print(" olmOCR-bench Results Summary (DeepSeek-OCR-2 via sglang)")
|
||||||
|
print(sep)
|
||||||
|
print(f"{'Split':<22} {'Tests':>8} {'Passed':>8} {'Score':>8}")
|
||||||
|
print("-" * 50)
|
||||||
|
splits = list(all_results.keys())
|
||||||
|
scores = []
|
||||||
|
for split in splits:
|
||||||
|
r = all_results[split]
|
||||||
|
score = r.get("overall_score", 0.0)
|
||||||
|
scores.append(score)
|
||||||
|
print(
|
||||||
|
f"{split:<22} {r['total_tests']:>8} {r['total_passed']:>8} {score:>7.1f}%"
|
||||||
|
)
|
||||||
|
print("-" * 50)
|
||||||
|
|
||||||
|
total_tests = sum(r["total_tests"] for r in all_results.values())
|
||||||
|
total_passed = sum(r["total_passed"] for r in all_results.values())
|
||||||
|
overall = round(100.0 * total_passed / total_tests, 1) if total_tests > 0 else 0.0
|
||||||
|
mean_score = round(sum(scores) / len(scores), 1) if scores else 0.0
|
||||||
|
print(f"{'TOTAL':<22} {total_tests:>8} {total_passed:>8} {overall:>7.1f}%")
|
||||||
|
print(f"{'Mean across splits':<22} {'':>17} {mean_score:>7.1f}%")
|
||||||
|
print(sep)
|
||||||
|
|
||||||
|
# Per-type breakdown
|
||||||
|
all_types: set = set()
|
||||||
|
for r in all_results.values():
|
||||||
|
all_types.update(r.get("by_type", {}).keys())
|
||||||
|
|
||||||
|
if all_types:
|
||||||
|
print("\nPer-test-type breakdown:")
|
||||||
|
print(f"{'Test Type':<35} {'Tests':>8} {'Score':>8}")
|
||||||
|
print("-" * 55)
|
||||||
|
for t in sorted(all_types):
|
||||||
|
totals = {"passed": 0, "total": 0}
|
||||||
|
for r in all_results.values():
|
||||||
|
counts = r.get("by_type_counts", {}).get(t, {"passed": 0, "total": 0})
|
||||||
|
totals["passed"] += counts["passed"]
|
||||||
|
totals["total"] += counts["total"]
|
||||||
|
type_score = (
|
||||||
|
round(100.0 * totals["passed"] / totals["total"], 1)
|
||||||
|
if totals["total"] > 0
|
||||||
|
else 0.0
|
||||||
|
)
|
||||||
|
print(f"{t:<35} {totals['total']:>8} {type_score:>7.1f}%")
|
||||||
|
print(sep)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Normalized Edit Distance (OmniDocBench-style text quality metric) ─────────
|
||||||
|
|
||||||
|
|
||||||
|
def normalized_edit_distance(pred: str, ref: str) -> float:
|
||||||
|
"""
|
||||||
|
Character-level Normalized Edit Distance in [0, 1].
|
||||||
|
0.0 = identical, 1.0 = completely different.
|
||||||
|
"""
|
||||||
|
pred = normalize_text(pred.strip())
|
||||||
|
ref = normalize_text(ref.strip())
|
||||||
|
if not ref and not pred:
|
||||||
|
return 0.0
|
||||||
|
if not ref or not pred:
|
||||||
|
return 1.0
|
||||||
|
|
||||||
|
m, n = len(pred), len(ref)
|
||||||
|
# Space-optimised single-row DP
|
||||||
|
dp = list(range(n + 1))
|
||||||
|
for i in range(1, m + 1):
|
||||||
|
prev = dp[0]
|
||||||
|
dp[0] = i
|
||||||
|
for j in range(1, n + 1):
|
||||||
|
temp = dp[j]
|
||||||
|
if pred[i - 1] == ref[j - 1]:
|
||||||
|
dp[j] = prev
|
||||||
|
else:
|
||||||
|
dp[j] = 1 + min(prev, dp[j], dp[j - 1])
|
||||||
|
prev = temp
|
||||||
|
|
||||||
|
return dp[n] / max(m, n)
|
||||||
@@ -0,0 +1,381 @@
|
|||||||
|
"""
|
||||||
|
Generate a self-contained HTML verification report from olmOCR-bench results.
|
||||||
|
|
||||||
|
Requires results saved with --save-raw-outputs.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
# 1. Run benchmark with raw outputs saved
|
||||||
|
python benchmark/ocr/bench_sglang.py --port 30000 --split arxiv_math \\
|
||||||
|
--max-samples 20 --save-raw-outputs
|
||||||
|
|
||||||
|
# 2. Generate HTML report for a single split
|
||||||
|
python benchmark/ocr/generate_report.py --split arxiv_math
|
||||||
|
|
||||||
|
# 3. Generate HTML report for all splits in a results directory
|
||||||
|
python benchmark/ocr/generate_report.py --results-dir ./ocr_bench_results
|
||||||
|
|
||||||
|
# 4. Show only failing tests
|
||||||
|
python benchmark/ocr/generate_report.py --split arxiv_math --failures-only
|
||||||
|
|
||||||
|
Open the generated .html file in any browser — formulas are rendered via MathJax.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import html
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_MATH_DELIM_RE = re.compile(
|
||||||
|
r"(\\\[[\s\S]*?\\\]" # \[...\]
|
||||||
|
r"|\\\([\s\S]*?\\\)" # \(...\)
|
||||||
|
r"|\$\$[\s\S]*?\$\$" # $$...$$
|
||||||
|
r"|\$[^$\n]+?\$)", # $...$
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _ocr_to_html(text: str) -> str:
|
||||||
|
"""
|
||||||
|
Convert raw OCR output to readable HTML.
|
||||||
|
Preserves LaTeX delimiters for MathJax, strips bounding-box annotations,
|
||||||
|
and wraps block-level elements in <p> tags.
|
||||||
|
"""
|
||||||
|
lines = text.splitlines()
|
||||||
|
out_lines = []
|
||||||
|
for line in lines:
|
||||||
|
# Strip bounding box annotations like text[[x1, y1, x2, y2]]
|
||||||
|
line = re.sub(r"^\s*\w[\w_]*\[\[\d+,\s*\d+,\s*\d+,\s*\d+\]\]\s*", "", line)
|
||||||
|
# HTML-escape everything EXCEPT LaTeX delimiters
|
||||||
|
parts = _MATH_DELIM_RE.split(line)
|
||||||
|
escaped = ""
|
||||||
|
for i, part in enumerate(parts):
|
||||||
|
if i % 2 == 0:
|
||||||
|
escaped += html.escape(part)
|
||||||
|
else:
|
||||||
|
escaped += part # LaTeX — pass through for MathJax
|
||||||
|
out_lines.append(escaped)
|
||||||
|
return "<br>".join(out_lines)
|
||||||
|
|
||||||
|
|
||||||
|
def _latex_to_display(latex: str) -> str:
|
||||||
|
"""Wrap a raw LaTeX string in display-math delimiters for MathJax."""
|
||||||
|
stripped = latex.strip()
|
||||||
|
# Already delimited → pass through
|
||||||
|
if stripped.startswith(("\\[", "$$", "\\(")):
|
||||||
|
return stripped
|
||||||
|
return f"\\[ {stripped} \\]"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# HTML template
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_HTML_HEAD = """\
|
||||||
|
<!DOCTYPE html>
|
||||||
|
<html lang="en">
|
||||||
|
<head>
|
||||||
|
<meta charset="UTF-8">
|
||||||
|
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||||
|
<title>{title}</title>
|
||||||
|
<script>
|
||||||
|
window.MathJax = {{
|
||||||
|
tex: {{ inlineMath: [['\\\\(','\\\\)'], ['$','$']], displayMath: [['\\\\[','\\\\]'],['$$','$$']] }},
|
||||||
|
options: {{ skipHtmlTags: ['script','noscript','style','textarea'] }}
|
||||||
|
}};
|
||||||
|
</script>
|
||||||
|
<script src="https://cdn.jsdelivr.net/npm/mathjax@3/es5/tex-chtml.js" async></script>
|
||||||
|
<style>
|
||||||
|
body {{ font-family: 'Segoe UI', sans-serif; margin: 0; background: #f5f5f5; color: #222; }}
|
||||||
|
h1 {{ background: #1a73e8; color: #fff; margin: 0; padding: 16px 24px; font-size: 1.3em; }}
|
||||||
|
.summary {{ background: #fff; border-bottom: 1px solid #ddd; padding: 12px 24px;
|
||||||
|
display: flex; gap: 32px; flex-wrap: wrap; font-size: 0.9em; }}
|
||||||
|
.stat {{ display: flex; flex-direction: column; }}
|
||||||
|
.stat span:first-child {{ font-weight: 600; font-size: 1.1em; }}
|
||||||
|
.sample {{ background: #fff; border: 1px solid #ddd; border-radius: 6px;
|
||||||
|
margin: 16px 24px; overflow: hidden; }}
|
||||||
|
.sample-header {{ padding: 10px 16px; font-weight: 600; font-size: 0.85em;
|
||||||
|
background: #f0f0f0; display: flex; justify-content: space-between;
|
||||||
|
align-items: center; cursor: pointer; user-select: none; }}
|
||||||
|
.sample-header:hover {{ background: #e8e8e8; }}
|
||||||
|
.sample-body {{ padding: 0 16px 16px; }}
|
||||||
|
.tests {{ margin-top: 12px; display: flex; flex-direction: column; gap: 12px; }}
|
||||||
|
.test {{ border: 1px solid #ddd; border-radius: 4px; overflow: hidden; }}
|
||||||
|
.test-header {{ padding: 6px 12px; font-size: 0.8em; font-weight: 600;
|
||||||
|
display: flex; gap: 12px; align-items: center; }}
|
||||||
|
.pass {{ background: #e6f4ea; border-left: 4px solid #34a853; }}
|
||||||
|
.fail {{ background: #fce8e6; border-left: 4px solid #ea4335; }}
|
||||||
|
.test-body {{ display: grid; grid-template-columns: 1fr 1fr; gap: 0; }}
|
||||||
|
.pane {{ padding: 10px 14px; font-size: 0.85em; }}
|
||||||
|
.pane:first-child {{ border-right: 1px solid #ddd; }}
|
||||||
|
.pane h4 {{ margin: 0 0 6px; font-size: 0.75em; text-transform: uppercase;
|
||||||
|
letter-spacing: 0.05em; color: #666; }}
|
||||||
|
.pane pre {{ margin: 0; white-space: pre-wrap; word-break: break-all;
|
||||||
|
font-family: 'Cascadia Code', 'Fira Code', monospace; font-size: 0.9em;
|
||||||
|
background: #f8f8f8; padding: 6px 8px; border-radius: 3px; }}
|
||||||
|
.rendered {{ font-size: 1em; padding: 4px 0; min-height: 24px; }}
|
||||||
|
.ocr-output {{ grid-column: 1 / -1; padding: 10px 14px; font-size: 0.82em; }}
|
||||||
|
.ocr-output h4 {{ margin: 0 0 6px; font-size: 0.75em; text-transform: uppercase;
|
||||||
|
letter-spacing: 0.05em; color: #666; }}
|
||||||
|
.ocr-raw {{ white-space: pre-wrap; word-break: break-all; max-height: 300px;
|
||||||
|
overflow-y: auto; background: #f8f8f8; padding: 8px 10px; border-radius: 3px;
|
||||||
|
font-family: monospace; font-size: 0.85em; line-height: 1.5; }}
|
||||||
|
.badge {{ display: inline-block; padding: 2px 8px; border-radius: 10px;
|
||||||
|
font-size: 0.78em; font-weight: 700; }}
|
||||||
|
.badge-pass {{ background: #34a853; color: #fff; }}
|
||||||
|
.badge-fail {{ background: #ea4335; color: #fff; }}
|
||||||
|
.badge-score {{ background: #1a73e8; color: #fff; }}
|
||||||
|
.toggle {{ font-size: 0.8em; color: #666; }}
|
||||||
|
details summary {{ list-style: none; }}
|
||||||
|
details summary::-webkit-details-marker {{ display: none; }}
|
||||||
|
</style>
|
||||||
|
</head>
|
||||||
|
<body>
|
||||||
|
<h1>{title}</h1>
|
||||||
|
"""
|
||||||
|
|
||||||
|
_HTML_TAIL = """\
|
||||||
|
</body>
|
||||||
|
</html>
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _render_sample(sample: dict, failures_only: bool) -> str:
|
||||||
|
"""Render one sample (PDF page) as an HTML block."""
|
||||||
|
pdf = sample.get("pdf", sample.get("pdf_rel", "?"))
|
||||||
|
page = sample.get("page", "?")
|
||||||
|
passed = sample.get("passed", 0)
|
||||||
|
total = sample.get("total", 0)
|
||||||
|
error = sample.get("error")
|
||||||
|
ocr_output = sample.get("ocr_output", "")
|
||||||
|
test_results = sample.get("test_results", [])
|
||||||
|
test_inputs = sample.get("test_inputs", [])
|
||||||
|
|
||||||
|
if failures_only and passed == total and not error:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
pct = f"{100*passed//total}%" if total else "—"
|
||||||
|
header_cls = "fail" if (error or passed < total) else "pass"
|
||||||
|
|
||||||
|
parts = [f'<div class="sample">']
|
||||||
|
parts.append(
|
||||||
|
f'<details {"open" if (error or passed < total) else ""}>'
|
||||||
|
f'<summary class="sample-header {header_cls}">'
|
||||||
|
f"<span>📄 {html.escape(pdf)} · page {page}</span>"
|
||||||
|
f"<span>"
|
||||||
|
f'<span class="badge badge-score">{passed}/{total} {pct}</span>'
|
||||||
|
f"</span>"
|
||||||
|
f"</summary>"
|
||||||
|
)
|
||||||
|
parts.append('<div class="sample-body">')
|
||||||
|
|
||||||
|
if error:
|
||||||
|
parts.append(
|
||||||
|
f'<p style="color:#ea4335;font-weight:600">ERROR: {html.escape(str(error))}</p>'
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parts.append('<div class="tests">')
|
||||||
|
for i, tr in enumerate(test_results):
|
||||||
|
if failures_only and tr.get("passed"):
|
||||||
|
continue
|
||||||
|
ttype = tr.get("type", "?")
|
||||||
|
ok = tr.get("passed", False)
|
||||||
|
badge = (
|
||||||
|
'<span class="badge badge-pass">PASS</span>'
|
||||||
|
if ok
|
||||||
|
else '<span class="badge badge-fail">FAIL</span>'
|
||||||
|
)
|
||||||
|
test_err = tr.get("error", "")
|
||||||
|
# Get expected values from test_inputs (parallel list)
|
||||||
|
ti = test_inputs[i] if i < len(test_inputs) else {}
|
||||||
|
|
||||||
|
test_cls = "pass" if ok else "fail"
|
||||||
|
parts.append(f'<div class="test">')
|
||||||
|
parts.append(
|
||||||
|
f'<div class="test-header {test_cls}">'
|
||||||
|
f"{badge} <code>type={ttype!r}</code>"
|
||||||
|
+ (
|
||||||
|
f' <span style="color:#c00">{html.escape(test_err)}</span>'
|
||||||
|
if test_err
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
+ f"</div>"
|
||||||
|
)
|
||||||
|
# Expected pane
|
||||||
|
parts.append('<div class="test-body">')
|
||||||
|
parts.append('<div class="pane">')
|
||||||
|
parts.append("<h4>Expected</h4>")
|
||||||
|
if ttype in ("math", "math_formula_accuracy"):
|
||||||
|
latex = ti.get("math", "")
|
||||||
|
parts.append(f"<pre>{html.escape(latex)}</pre>")
|
||||||
|
if latex:
|
||||||
|
parts.append(
|
||||||
|
f'<div class="rendered">{_latex_to_display(latex)}</div>'
|
||||||
|
)
|
||||||
|
elif ttype in ("present", "absent", "text_presence", "text_absence"):
|
||||||
|
parts.append(f'<pre>{html.escape(ti.get("text", ""))}</pre>')
|
||||||
|
elif ttype in ("order", "natural_reading_order"):
|
||||||
|
before = ti.get("before", "")
|
||||||
|
after = ti.get("after", "")
|
||||||
|
parts.append(
|
||||||
|
f"<pre>before: {html.escape(before)}\nafter: {html.escape(after)}</pre>"
|
||||||
|
)
|
||||||
|
parts.append("</div>")
|
||||||
|
# Right pane — OCR formulas extracted (if math type)
|
||||||
|
parts.append('<div class="pane">')
|
||||||
|
parts.append("<h4>OCR extracted formulas</h4>")
|
||||||
|
if ttype in ("math", "math_formula_accuracy") and ocr_output:
|
||||||
|
matches = _MATH_DELIM_RE.findall(ocr_output)
|
||||||
|
if matches:
|
||||||
|
for m in matches[:6]: # show at most 6 matches
|
||||||
|
parts.append(f'<div class="rendered">{m}</div>')
|
||||||
|
if len(matches) > 6:
|
||||||
|
parts.append(
|
||||||
|
f'<p style="color:#888;font-size:0.8em">… and {len(matches)-6} more</p>'
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parts.append(
|
||||||
|
'<p style="color:#888;font-size:0.8em">No LaTeX delimiters found in OCR output.</p>'
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parts.append(
|
||||||
|
'<p style="color:#888;font-size:0.8em">(see full OCR output below)</p>'
|
||||||
|
)
|
||||||
|
parts.append("</div>")
|
||||||
|
parts.append("</div>") # test-body
|
||||||
|
parts.append("</div>") # test
|
||||||
|
|
||||||
|
parts.append("</div>") # tests
|
||||||
|
|
||||||
|
# Full OCR output
|
||||||
|
if ocr_output:
|
||||||
|
parts.append('<div class="ocr-output">')
|
||||||
|
parts.append(f"<h4>Full OCR output ({len(ocr_output)} chars)</h4>")
|
||||||
|
parts.append(f'<div class="ocr-raw">{_ocr_to_html(ocr_output)}</div>')
|
||||||
|
parts.append("</div>")
|
||||||
|
elif not error:
|
||||||
|
parts.append(
|
||||||
|
'<p style="color:#888;font-size:0.85em">No raw OCR output stored. '
|
||||||
|
"Re-run with <code>--save-raw-outputs</code> to include it here.</p>"
|
||||||
|
)
|
||||||
|
|
||||||
|
parts.append("</div>") # sample-body
|
||||||
|
parts.append("</details>")
|
||||||
|
parts.append("</div>") # sample
|
||||||
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Report generation
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def generate_report(
|
||||||
|
split_name: str,
|
||||||
|
split_result: dict,
|
||||||
|
out_path: Path,
|
||||||
|
failures_only: bool = False,
|
||||||
|
) -> None:
|
||||||
|
total_passed = split_result.get("total_passed", 0)
|
||||||
|
total_tests = split_result.get("total_tests", 0)
|
||||||
|
overall = split_result.get("overall_score", 0.0)
|
||||||
|
error_samples = split_result.get("error_samples", 0)
|
||||||
|
total_samples = split_result.get("total_samples", 0)
|
||||||
|
|
||||||
|
title = f"OCR Bench — {split_name}"
|
||||||
|
body = _HTML_HEAD.format(title=title)
|
||||||
|
|
||||||
|
# Summary bar
|
||||||
|
by_type = split_result.get("by_type", {})
|
||||||
|
type_badges = " ".join(
|
||||||
|
f'<span class="badge badge-score">{t}: {v:.1f}%</span>'
|
||||||
|
for t, v in by_type.items()
|
||||||
|
)
|
||||||
|
body += f"""
|
||||||
|
<div class="summary">
|
||||||
|
<div class="stat"><span>{total_passed}/{total_tests}</span><span>tests passed</span></div>
|
||||||
|
<div class="stat"><span>{overall:.1f}%</span><span>overall score</span></div>
|
||||||
|
<div class="stat"><span>{error_samples}/{total_samples}</span><span>error samples</span></div>
|
||||||
|
<div class="stat">{type_badges}</div>
|
||||||
|
</div>
|
||||||
|
"""
|
||||||
|
if failures_only:
|
||||||
|
body += '<p style="margin:12px 24px;color:#888;font-size:0.9em">Showing failures only.</p>'
|
||||||
|
|
||||||
|
samples = split_result.get("samples", [])
|
||||||
|
if not samples:
|
||||||
|
body += '<p style="margin:24px;color:#888">No per-sample data found. The result JSON may not include sample-level details.</p>'
|
||||||
|
else:
|
||||||
|
for sample in samples:
|
||||||
|
body += _render_sample(sample, failures_only)
|
||||||
|
|
||||||
|
body += _HTML_TAIL
|
||||||
|
out_path.write_text(body, encoding="utf-8")
|
||||||
|
print(f"Report → {out_path}")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Entry point
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Generate HTML verification report from olmOCR-bench results",
|
||||||
|
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--results-dir",
|
||||||
|
default="./ocr_bench_results",
|
||||||
|
help="Directory containing result JSON files",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--split",
|
||||||
|
default=None,
|
||||||
|
help="Single split to report (default: all JSONs in results-dir)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output-dir",
|
||||||
|
default=None,
|
||||||
|
help="Where to write HTML files (default: same as results-dir)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--failures-only",
|
||||||
|
action="store_true",
|
||||||
|
help="Include only samples/tests that failed",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
results_dir = Path(args.results_dir)
|
||||||
|
output_dir = Path(args.output_dir) if args.output_dir else results_dir
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
if args.split:
|
||||||
|
json_files = [results_dir / f"{args.split}.json"]
|
||||||
|
else:
|
||||||
|
json_files = sorted(results_dir.glob("*.json"))
|
||||||
|
json_files = [f for f in json_files if f.stem != "summary"]
|
||||||
|
|
||||||
|
if not json_files:
|
||||||
|
sys.exit(f"No result JSON files found in {results_dir}")
|
||||||
|
|
||||||
|
for jf in json_files:
|
||||||
|
if not jf.exists():
|
||||||
|
print(f" SKIP (not found): {jf}")
|
||||||
|
continue
|
||||||
|
split_name = jf.stem
|
||||||
|
with open(jf, encoding="utf-8") as fh:
|
||||||
|
data = json.load(fh)
|
||||||
|
|
||||||
|
suffix = "_failures" if args.failures_only else ""
|
||||||
|
out_path = output_dir / f"{split_name}{suffix}_report.html"
|
||||||
|
generate_report(split_name, data, out_path, failures_only=args.failures_only)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -142,6 +142,7 @@ test = [
|
|||||||
"addict",
|
"addict",
|
||||||
"auto-round>=0.13.1",
|
"auto-round>=0.13.1",
|
||||||
"bitsandbytes",
|
"bitsandbytes",
|
||||||
|
"pymupdf",
|
||||||
"diff-cover",
|
"diff-cover",
|
||||||
"expecttest",
|
"expecttest",
|
||||||
"granian>=2.6.0",
|
"granian>=2.6.0",
|
||||||
|
|||||||
@@ -102,6 +102,7 @@ tracing = [
|
|||||||
]
|
]
|
||||||
test = [
|
test = [
|
||||||
"accelerate",
|
"accelerate",
|
||||||
|
"pymupdf",
|
||||||
"expecttest",
|
"expecttest",
|
||||||
"jsonlines",
|
"jsonlines",
|
||||||
"matplotlib",
|
"matplotlib",
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ tracing = [
|
|||||||
|
|
||||||
test = [
|
test = [
|
||||||
"accelerate",
|
"accelerate",
|
||||||
|
"pymupdf",
|
||||||
"expecttest",
|
"expecttest",
|
||||||
"gguf",
|
"gguf",
|
||||||
"jsonlines",
|
"jsonlines",
|
||||||
|
|||||||
@@ -162,6 +162,7 @@ diffusion_mps = [
|
|||||||
|
|
||||||
test = [
|
test = [
|
||||||
"accelerate",
|
"accelerate",
|
||||||
|
"pymupdf",
|
||||||
"expecttest",
|
"expecttest",
|
||||||
"gguf",
|
"gguf",
|
||||||
"jsonlines",
|
"jsonlines",
|
||||||
|
|||||||
@@ -100,6 +100,7 @@ tracing = [
|
|||||||
test = [
|
test = [
|
||||||
"accelerate",
|
"accelerate",
|
||||||
"bitsandbytes",
|
"bitsandbytes",
|
||||||
|
"pymupdf",
|
||||||
"expecttest",
|
"expecttest",
|
||||||
"jsonlines",
|
"jsonlines",
|
||||||
"lm-eval[api]>=0.4.9.2",
|
"lm-eval[api]>=0.4.9.2",
|
||||||
|
|||||||
Reference in New Issue
Block a user