72 lines
2.8 KiB
Python
72 lines
2.8 KiB
Python
"""Download official model weights and verify them on a synthetic page."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import os
|
|
from pathlib import Path
|
|
import tempfile
|
|
|
|
from PIL import Image, ImageDraw
|
|
|
|
|
|
ROOT = Path(os.environ.get("OCR_MODEL_HOME", Path(__file__).resolve().parents[1] / ".cache")).expanduser().resolve()
|
|
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--classic", action="store_true")
|
|
parser.add_argument("--vision", action="store_true")
|
|
parser.add_argument("--device", default="cpu", help="Vision device, for example cpu or gpu:0")
|
|
args = parser.parse_args()
|
|
if not args.classic and not args.vision:
|
|
parser.error("choose --classic and/or --vision")
|
|
|
|
ROOT.mkdir(parents=True, exist_ok=True)
|
|
os.environ["PADDLE_PDX_CACHE_HOME"] = str(ROOT)
|
|
os.environ["HF_HOME"] = str(ROOT / "hf")
|
|
os.environ.pop("HF_HUB_OFFLINE", None)
|
|
os.environ.pop("TRANSFORMERS_OFFLINE", None)
|
|
with tempfile.TemporaryDirectory(prefix="ocr-synthetic-setup-") as temporary:
|
|
image = Path(temporary) / "synthetic.png"
|
|
page = Image.new("RGB", (1000, 600), "white")
|
|
draw = ImageDraw.Draw(page)
|
|
draw.text((50, 50), "SYNTHETIC TEST / NO PERSONAL DATA", fill="black")
|
|
draw.text((50, 130), "D.3 SAMPLE 123", fill="black")
|
|
page.save(image)
|
|
if args.classic:
|
|
from paddleocr import PaddleOCR
|
|
engine = PaddleOCR(
|
|
text_detection_model_name="PP-OCRv5_mobile_det",
|
|
text_recognition_model_name="latin_PP-OCRv5_mobile_rec",
|
|
use_doc_orientation_classify=False,
|
|
use_doc_unwarping=False,
|
|
use_textline_orientation=False,
|
|
device="cpu",
|
|
)
|
|
list(engine.predict(str(image)))
|
|
if args.vision:
|
|
from paddleocr import PaddleOCRVL
|
|
engine = PaddleOCRVL(
|
|
pipeline_version="v1.6", use_layout_detection=True,
|
|
use_doc_orientation_classify=False, use_doc_unwarping=False,
|
|
device=args.device,
|
|
)
|
|
list(engine.predict(str(image)))
|
|
|
|
models = ROOT / "official_models"
|
|
checks = []
|
|
if args.classic:
|
|
checks.extend((models / name / "inference.pdiparams" for name in
|
|
("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec")))
|
|
if args.vision:
|
|
checks.extend((models / "PP-DocLayoutV3" / "inference.pdiparams",
|
|
models / "PaddleOCR-VL-1.6" / "model.safetensors"))
|
|
if any(not path.is_file() for path in checks):
|
|
raise SystemExit("Official model download did not create all expected files")
|
|
print("Official models ready; only synthetic pixels were used")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|