Bina-0.2-Rizeh / bina_page_ocr.py
Reza2kn's picture
Complete self-contained Bina 0.2 Rizeh release and quick start
2af6ae7 verified
Raw
History Blame Contribute Delete
5.28 kB
#!/usr/bin/env python3
"""Offline full-page Persian OCR using bundled detector plus a Bina recognizer."""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
from typing import Any, Iterator
from paddleocr import PaddleOCR
_LTR_RUN = re.compile(r"[a-zA-Z0-9 :*./%+-]")
def pred_reverse(text: str) -> str:
segments: list[str] = []
current_ltr = ""
for character in text:
if _LTR_RUN.search(character):
current_ltr += character
continue
if current_ltr:
segments.append(current_ltr)
current_ltr = ""
segments.append(character)
if current_ltr:
segments.append(current_ltr)
return "".join(reversed(segments))
def _json_payload(result: Any) -> dict[str, Any]:
payload = result.json() if callable(result.json) else result.json
return payload.get("res", payload)
def _box_list(box: Any) -> list[float]:
if hasattr(box, "tolist"):
box = box.tolist()
return [float(value) for value in box]
def ordered_lines(payload: dict[str, Any]) -> list[dict[str, Any]]:
visual_texts = [str(value) for value in payload.get("rec_texts", [])]
texts = [pred_reverse(value) for value in visual_texts]
scores = [float(value) for value in payload.get("rec_scores", [0.0] * len(texts))]
boxes = payload.get("rec_boxes", [])
if len(boxes) != len(texts):
return [{"text": text, "score": scores[i], "raw_visual_text": visual_texts[i]} for i, text in enumerate(texts)]
items = []
for i, (text, box) in enumerate(zip(texts, boxes)):
x0, y0, x1, y1 = _box_list(box)
items.append({
"text": text,
"score": scores[i],
"raw_visual_text": visual_texts[i],
"box": [x0, y0, x1, y1],
"x": (x0 + x1) / 2,
"y": (y0 + y1) / 2,
"height": max(y1 - y0, 1.0),
})
items.sort(key=lambda item: item["y"])
rows: list[list[dict[str, Any]]] = []
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])
output = []
for row_index, row in enumerate(rows):
row.sort(key=lambda item: item["x"], reverse=True)
for item in row:
output.append({key: value for key, value in item.items() if key not in {"x", "y", "height"}} | {"row": row_index})
return output
class BinaPageOCR:
def __init__(self, model_dir: str | Path | None = None, detector_dir: str | Path | None = None,
device: str | None = None, score_threshold: float = 0.0) -> None:
base = Path(__file__).resolve().parent
model_dir = Path(model_dir) if model_dir else base / "inference"
detector_dir = Path(detector_dir) if detector_dir else base / "detector"
options: dict[str, Any] = {
"text_detection_model_dir": str(detector_dir),
"text_recognition_model_dir": str(model_dir),
"use_doc_orientation_classify": False,
"use_doc_unwarping": False,
"use_textline_orientation": False,
"text_rec_score_thresh": score_threshold,
}
if device:
options["device"] = device
self._ocr = PaddleOCR(**options)
def predict(self, inputs: list[str | Path]) -> Iterator[dict[str, Any]]:
for source in inputs:
for page_index, result in enumerate(self._ocr.predict(str(source))):
raw = _json_payload(result)
lines = ordered_lines(raw)
row_text: dict[int, list[str]] = {}
for line in lines:
row_text.setdefault(int(line.get("row", len(row_text))), []).append(line["text"])
text = "\n".join(" ".join(row_text[index]).strip() for index in sorted(row_text) if row_text[index])
yield {"input_path": str(source), "page_index": page_index, "text": text, "lines": lines}
def main() -> int:
parser = argparse.ArgumentParser(description="Run self-contained full-page Persian OCR with Bina 0.2")
parser.add_argument("inputs", nargs="+")
parser.add_argument("--model-dir", default=str(Path(__file__).resolve().parent / "inference"))
parser.add_argument("--detector-dir", default=str(Path(__file__).resolve().parent / "detector"))
parser.add_argument("--device", default="cpu", help="cpu, gpu:0, ...")
parser.add_argument("--score-threshold", type=float, default=0.0)
parser.add_argument("--output", type=Path)
args = parser.parse_args()
model = BinaPageOCR(args.model_dir, args.detector_dir, args.device, args.score_threshold)
predictions = list(model.predict(args.inputs))
for prediction in predictions:
print(json.dumps(prediction, ensure_ascii=False))
if args.output:
args.output.write_text(json.dumps(predictions, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
return 0
if __name__ == "__main__":
raise SystemExit(main())