File size: 4,889 Bytes
054400d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | #!/usr/bin/env python3
"""Prepare a small Persian PP-OCRv6 recognition pilot on a remote GPU."""
from __future__ import annotations
import argparse
import io
import json
import unicodedata
from collections import Counter
from pathlib import Path
import pyarrow.parquet as pq
from huggingface_hub import HfApi, hf_hub_download
from PIL import Image
def clean_text(value: str) -> str:
value = unicodedata.normalize("NFC", value).replace("\t", " ")
return " ".join(value.split())
def image_bytes(value: object) -> bytes:
if isinstance(value, dict) and value.get("bytes") is not None:
return value["bytes"]
raise ValueError("Expected an embedded Hugging Face image payload")
def read_gt(path: Path) -> list[tuple[Path, str]]:
rows = []
for raw in path.read_text(encoding="utf-8").splitlines():
key, text = raw.split(maxsplit=1)
rows.append((path.parent / f"{key}.png", clean_text(text)))
return rows
def first_parquet(repo: str) -> str:
files = HfApi().list_repo_files(repo, repo_type="dataset")
candidates = sorted(x for x in files if x.endswith(".parquet") and x.startswith("data/"))
if not candidates:
raise RuntimeError(f"No data parquet found in {repo}")
return candidates[0]
def extract_printed(repo: str, output: Path, train_count: int, eval_count: int) -> tuple[list, list]:
source = hf_hub_download(
repo, first_parquet(repo), repo_type="dataset", local_dir=output / "hf"
)
table = pq.read_table(source, columns=["image", "label"]).slice(0, train_count + eval_count)
rows = table.to_pylist()
if len(rows) < train_count + eval_count:
raise RuntimeError("Printed pilot shard does not contain enough rows")
destination = output / "images" / "printed"
destination.mkdir(parents=True, exist_ok=True)
converted = []
for index, row in enumerate(rows):
path = destination / f"{index:07d}.jpg"
with Image.open(io.BytesIO(image_bytes(row["image"]))) as image:
image.convert("RGB").save(path, "JPEG", quality=95)
converted.append((path, clean_text(row["label"])))
return converted[:train_count], converted[train_count:]
def load_hardword_text(repo: str, output: Path) -> list[str]:
source = hf_hub_download(
repo, first_parquet(repo), repo_type="dataset", local_dir=output / "hf"
)
return [clean_text(x) for x in pq.read_table(source, columns=["sentence"])["sentence"].to_pylist()]
def write_list(path: Path, rows: list[tuple[Path, str]], root: Path) -> None:
with path.open("w", encoding="utf-8") as handle:
for image, text in rows:
handle.write(f"{image.relative_to(root)}\t{text}\n")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--phtd", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--printed-train", type=int, default=4000)
parser.add_argument("--printed-eval", type=int, default=500)
args = parser.parse_args()
args.output.mkdir(parents=True, exist_ok=True)
handwriting_train = read_gt(args.phtd / "train" / "gt.txt")
handwriting_eval = read_gt(args.phtd / "val" / "gt.txt")
printed_train, printed_eval = extract_printed(
"Reza2kn/persian-printed-ocr-3.5m",
args.output,
args.printed_train,
args.printed_eval,
)
hardwords = load_hardword_text("Reza2kn/visualears-hardword-sentences", args.output)
train = handwriting_train + printed_train
evaluation = handwriting_eval + printed_eval
data_root = args.output.parent
write_list(args.output / "train_list.txt", train, data_root)
write_list(args.output / "val_list.txt", evaluation, data_root)
frequencies = Counter("".join(text for _, text in train + evaluation) + "".join(hardwords))
characters = sorted(
char for char in frequencies
if char != " " and char.isprintable() and char not in "\r\n\t"
)
(args.output / "persian_dict.txt").write_text(
"".join(f"{char}\n" for char in characters), encoding="utf-8"
)
summary = {
"train_rows": len(train),
"eval_rows": len(evaluation),
"handwriting_train": len(handwriting_train),
"handwriting_eval": len(handwriting_eval),
"printed_train": len(printed_train),
"printed_eval": len(printed_eval),
"hardword_sentences_for_charset": len(hardwords),
"characters": len(characters),
"max_train_characters": max(map(lambda row: len(row[1]), train)),
"max_eval_characters": max(map(lambda row: len(row[1]), evaluation)),
}
(args.output / "summary.json").write_text(
json.dumps(summary, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
print(json.dumps(summary, ensure_ascii=False))
if __name__ == "__main__":
main()
|