#!/usr/bin/env python3 """Eight-way resumable printed-only exact benchmark for the Stage7 checkpoint.""" from __future__ import annotations import argparse import fcntl import json import os import re import time import traceback from pathlib import Path from typing import Any import cv2 import numpy as np import pyarrow.parquet as pq from paddleocr import PaddleOCR ROWS = 6_669 LTR_RUN = re.compile(r"[a-zA-Z0-9 :*./%+-]") def pred_reverse(text: str) -> str: segments = [] current = "" for character in text: if LTR_RUN.search(character): current += character else: if current: segments.append(current) current = "" segments.append(character) if current: segments.append(current) return "".join(reversed(segments)) def atomic_json(path: Path, payload: dict[str, Any]) -> None: temporary = path.with_suffix(path.suffix + f".{os.getpid()}.tmp") temporary.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n") os.replace(temporary, path) def ordered_text(payload: dict[str, Any]) -> tuple[str, int]: texts = [pred_reverse(str(text)) for text in payload.get("rec_texts", [])] boxes = payload.get("rec_boxes", []) if len(boxes) != len(texts): return "\n".join(text for text in texts if text.strip()), len(texts) items = [] for text, box in zip(texts, boxes): x0, y0, x1, y1 = (float(value) for value in box) items.append( { "text": text, "x": (x0 + x1) / 2, "y": (y0 + y1) / 2, "height": max(y1 - y0, 1), } ) items.sort(key=lambda item: item["y"]) rows = [] for item in items: if not rows: rows.append([item]) continue row = rows[-1] mean_y = sum(part["y"] for part in row) / len(row) mean_h = sum(part["height"] for part in row) / len(row) if abs(item["y"] - mean_y) <= 0.55 * max(item["height"], mean_h): row.append(item) else: rows.append([item]) lines = [] for row in rows: row.sort(key=lambda item: item["x"], reverse=True) text = " ".join(item["text"].strip() for item in row if item["text"].strip()) if text: lines.append(text) return "\n".join(lines), len(texts) def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--rank", type=int, required=True) parser.add_argument("--world-size", type=int, default=8) parser.add_argument("--data", type=Path, required=True) parser.add_argument("--model-dir", type=Path, required=True) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--device", default="gpu:0") args = parser.parse_args() args.output_dir.mkdir(parents=True, exist_ok=True) start = ROWS * args.rank // args.world_size stop = ROWS * (args.rank + 1) // args.world_size output_path = args.output_dir / f"predictions-rank{args.rank:02d}.jsonl" status_path = args.output_dir / f"status-rank{args.rank:02d}.json" done = set() if output_path.is_file(): for line in output_path.read_text(encoding="utf-8").splitlines(): if line.strip(): done.add(json.loads(line)["id"]) ocr = PaddleOCR( text_recognition_model_dir=str(args.model_dir), text_recognition_batch_size=64, use_doc_orientation_classify=False, use_doc_unwarping=False, use_textline_orientation=False, text_rec_score_thresh=0.0, text_detection_model_name="PP-OCRv6_medium_det", device=args.device, ) offset = 0 completed = 0 failures = 0 started = time.monotonic() with output_path.open("a", encoding="utf-8", buffering=1) as output: for shard in sorted((args.data / "data").glob("printed-*.parquet")): table = pq.read_table(shard, columns=["image", "label"]) for local, row in enumerate(table.to_pylist()): index = offset + local if not start <= index < stop or f"printed:{index}" in done: continue value = row["image"] encoded = value.get("bytes") if isinstance(value, dict) else value image = cv2.imdecode( np.frombuffer(encoded, dtype=np.uint8), cv2.IMREAD_COLOR, ) prediction = "" line_count = 0 error = None for attempt in range(1, 4): try: results = list(ocr.predict(image)) payload = results[0].json if callable(payload): payload = payload() prediction, line_count = ordered_text(payload.get("res", payload)) error = None break except Exception as exc: traceback.print_exc() error = f"{type(exc).__name__}: {exc}" time.sleep(attempt) failures += int(error is not None) result = { "id": f"printed:{index}", "split": "printed", "reference": str(row.get("label") or ""), "prediction": prediction, "detected_lines": line_count, } if error: result["error"] = error output.write(json.dumps(result, ensure_ascii=False) + "\n") completed += 1 offset += table.num_rows atomic_json( status_path, { "state": "complete", "rank": args.rank, "assigned": stop - start, "completed_this_run": completed, "failures": failures, "rate_pages_per_second": completed / max(time.monotonic() - started, 1e-9), }, ) with (args.output_dir / "merge.lock").open("a+") as lock: fcntl.flock(lock, fcntl.LOCK_EX) statuses = [] for rank in range(args.world_size): path = args.output_dir / f"status-rank{rank:02d}.json" if not path.is_file(): return statuses.append(json.loads(path.read_text())) rows = {} for rank in range(args.world_size): path = args.output_dir / f"predictions-rank{rank:02d}.jsonl" for line in path.read_text(encoding="utf-8").splitlines(): if line.strip(): row = json.loads(line) rows[row["id"]] = row expected = [f"printed:{index}" for index in range(ROWS)] if set(rows) != set(expected): raise RuntimeError("printed merge contract failed") with (args.output_dir / "predictions.jsonl.tmp").open("w", encoding="utf-8") as handle: for row_id in expected: handle.write(json.dumps(rows[row_id], ensure_ascii=False) + "\n") os.replace( args.output_dir / "predictions.jsonl.tmp", args.output_dir / "predictions.jsonl", ) atomic_json( args.output_dir / "run-summary.json", { "state": "complete", "rows": len(rows), "failures": sum(value["failures"] for value in statuses), }, ) if __name__ == "__main__": main()