Spaces:
Running
Running
Z User commited on
Commit ·
b0d0c0a
1
Parent(s): 1871500
feat: v4.4.0 — Training data improvements, checkpoint/resume, hot-reload, smart migrator
Browse files- Added checkpoint/resume to pdf_to_training_data.py for long PDFs
- Added hot_reload() and predict() to finetuning.py for immediate model use
- Added SmartMigrator utility for importing from legacy projects
- Added GRAM_006 translation rule (وفقاً ل)
- Enhanced Training Data tab with augmentation, val split, OCR toggle, gallery
- Updated docstring to list all 8 tabs
- Updated version to v4.4.0
- README.md +2 -2
- data/translation_rules.json +16 -0
- hf_app.py +41 -10
- modules/core/smart_migrator.py +351 -0
- modules/vision/finetuning.py +63 -0
- modules/vision/pdf_to_training_data.py +101 -2
README.md
CHANGED
|
@@ -11,7 +11,7 @@ license: mit
|
|
| 11 |
|
| 12 |
<div align="center">
|
| 13 |
|
| 14 |
-
# 🧠 OmniFile AI Processor v4.
|
| 15 |
|
| 16 |
**نظام ذكاء اصطناعي متكامل لمعالجة الملفات والنصوص والخط اليدوي**
|
| 17 |
**A Comprehensive AI System for File Processing, Text Analysis & Handwriting Recognition**
|
|
@@ -23,7 +23,7 @@ license: mit
|
|
| 23 |
[](https://github.com/DrAbdulmalek/OmniFile_Processor)
|
| 24 |
|
| 25 |
<p>
|
| 26 |
-
<b>Version:</b> v4.
|
| 27 |
<b>Status:</b> ✅ CI-Verified
|
| 28 |
</p>
|
| 29 |
|
|
|
|
| 11 |
|
| 12 |
<div align="center">
|
| 13 |
|
| 14 |
+
# 🧠 OmniFile AI Processor v4.4.0
|
| 15 |
|
| 16 |
**نظام ذكاء اصطناعي متكامل لمعالجة الملفات والنصوص والخط اليدوي**
|
| 17 |
**A Comprehensive AI System for File Processing, Text Analysis & Handwriting Recognition**
|
|
|
|
| 23 |
[](https://github.com/DrAbdulmalek/OmniFile_Processor)
|
| 24 |
|
| 25 |
<p>
|
| 26 |
+
<b>Version:</b> v4.4.0 |
|
| 27 |
<b>Status:</b> ✅ CI-Verified
|
| 28 |
</p>
|
| 29 |
|
data/translation_rules.json
CHANGED
|
@@ -97,6 +97,22 @@
|
|
| 97 |
"priority": 2,
|
| 98 |
"examples": []
|
| 99 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 100 |
{
|
| 101 |
"rule_id": "LEX_001",
|
| 102 |
"category": "lexical",
|
|
|
|
| 97 |
"priority": 2,
|
| 98 |
"examples": []
|
| 99 |
},
|
| 100 |
+
{
|
| 101 |
+
"rule_id": "GRAM_006",
|
| 102 |
+
"category": "grammatical",
|
| 103 |
+
"english_pattern": "according to",
|
| 104 |
+
"wrong_arabic": "طبقاً ل",
|
| 105 |
+
"correct_arabic": "وفقاً ل",
|
| 106 |
+
"rule_description": "استخدام 'وفقاً' بدلاً من 'طبقاً' للتعبير عن الإسناد",
|
| 107 |
+
"priority": 2,
|
| 108 |
+
"examples": [
|
| 109 |
+
{
|
| 110 |
+
"en": "According to the report",
|
| 111 |
+
"ar_wrong": "طبقاً للتقرير",
|
| 112 |
+
"ar_correct": "وفقاً للتقرير"
|
| 113 |
+
}
|
| 114 |
+
]
|
| 115 |
+
},
|
| 116 |
{
|
| 117 |
"rule_id": "LEX_001",
|
| 118 |
"category": "lexical",
|
hf_app.py
CHANGED
|
@@ -10,10 +10,11 @@ Tabs:
|
|
| 10 |
1. OCR Processing — EasyOCR + TrOCR ensemble
|
| 11 |
2. Text Correction — ar-corrector + pyspellchecker
|
| 12 |
3. PDF Processing — PyMuPDF (fitz) page-by-page extraction
|
| 13 |
-
4. Translation — Helsinki-NLP MarianMT models
|
| 14 |
-
5. Text Classification — Keyword-based
|
| 15 |
6. Evaluation — CER / WER metrics (Levenshtein + jiwer)
|
| 16 |
-
7.
|
|
|
|
| 17 |
|
| 18 |
Author: Dr Abdulmalek Tamer Al-husseini
|
| 19 |
Email: Abdulmalek.husseini@gmail.com
|
|
@@ -773,6 +774,9 @@ def generate_training_data(
|
|
| 773 |
pages: str = "1-5",
|
| 774 |
level: str = "word",
|
| 775 |
dpi: int = 300,
|
|
|
|
|
|
|
|
|
|
| 776 |
progress=gr.Progress(),
|
| 777 |
) -> Tuple[str, Optional[str]]:
|
| 778 |
"""
|
|
@@ -798,6 +802,8 @@ def generate_training_data(
|
|
| 798 |
dpi=dpi,
|
| 799 |
pages=pages,
|
| 800 |
max_image_height=64 if level == "character" else 0,
|
|
|
|
|
|
|
| 801 |
)
|
| 802 |
|
| 803 |
gen = TrainingDataGenerator(config=config)
|
|
@@ -805,11 +811,12 @@ def generate_training_data(
|
|
| 805 |
progress(0.3, desc="Loading OCR engine for text labels…")
|
| 806 |
# Try to get OCR engine for word labels
|
| 807 |
ocr_engine = None
|
| 808 |
-
|
| 809 |
-
|
| 810 |
-
|
| 811 |
-
|
| 812 |
-
|
|
|
|
| 813 |
|
| 814 |
progress(0.4, desc="Processing PDF pages…")
|
| 815 |
stats = gen.process_pdf(
|
|
@@ -847,7 +854,10 @@ def generate_training_data(
|
|
| 847 |
f"| **Train Samples** | {stats.get('train_samples', 0)} |\n"
|
| 848 |
f"| **Val Samples** | {stats.get('val_samples', 0)} |\n"
|
| 849 |
f"| **Level** | {level} |\n"
|
| 850 |
-
f"| **DPI** | {dpi} |\n
|
|
|
|
|
|
|
|
|
|
| 851 |
f"📥 Download the ZIP file containing:\n"
|
| 852 |
f"- `page_images/` — Full page renders\n"
|
| 853 |
f"- `word_crops/` — Individual word images\n"
|
|
@@ -1445,6 +1455,20 @@ def build_app() -> gr.Blocks:
|
|
| 1445 |
label="🔍 DPI",
|
| 1446 |
minimum=72, maximum=600, value=300, step=12,
|
| 1447 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1448 |
train_btn = gr.Button(
|
| 1449 |
"🚀 Generate Training Data / إنشاء بيانات التدريب",
|
| 1450 |
variant="primary", size="lg",
|
|
@@ -1452,6 +1476,12 @@ def build_app() -> gr.Blocks:
|
|
| 1452 |
|
| 1453 |
with gr.Column(scale=2):
|
| 1454 |
train_output = gr.Markdown("")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1455 |
train_files = gr.File(
|
| 1456 |
label="📥 Download Results / تنزيل النتائج",
|
| 1457 |
interactive=False,
|
|
@@ -1459,7 +1489,8 @@ def build_app() -> gr.Blocks:
|
|
| 1459 |
|
| 1460 |
train_btn.click(
|
| 1461 |
fn=generate_training_data,
|
| 1462 |
-
inputs=[train_file, train_pages, train_level, train_dpi
|
|
|
|
| 1463 |
outputs=[train_output, train_files],
|
| 1464 |
)
|
| 1465 |
|
|
|
|
| 10 |
1. OCR Processing — EasyOCR + TrOCR ensemble
|
| 11 |
2. Text Correction — ar-corrector + pyspellchecker
|
| 12 |
3. PDF Processing — PyMuPDF (fitz) page-by-page extraction
|
| 13 |
+
4. Translation — Helsinki-NLP MarianMT models + post-MT correction
|
| 14 |
+
5. Text Classification — Keyword-based multilingual categorization
|
| 15 |
6. Evaluation — CER / WER metrics (Levenshtein + jiwer)
|
| 16 |
+
7. Training Data — PDF to page/word/character crops for handwriting fine-tuning
|
| 17 |
+
8. About — Project info, links, author
|
| 18 |
|
| 19 |
Author: Dr Abdulmalek Tamer Al-husseini
|
| 20 |
Email: Abdulmalek.husseini@gmail.com
|
|
|
|
| 774 |
pages: str = "1-5",
|
| 775 |
level: str = "word",
|
| 776 |
dpi: int = 300,
|
| 777 |
+
enable_augmentation: bool = False,
|
| 778 |
+
val_split: float = 0.1,
|
| 779 |
+
enable_ocr: bool = True,
|
| 780 |
progress=gr.Progress(),
|
| 781 |
) -> Tuple[str, Optional[str]]:
|
| 782 |
"""
|
|
|
|
| 802 |
dpi=dpi,
|
| 803 |
pages=pages,
|
| 804 |
max_image_height=64 if level == "character" else 0,
|
| 805 |
+
enable_augmentation=enable_augmentation,
|
| 806 |
+
val_ratio=val_split,
|
| 807 |
)
|
| 808 |
|
| 809 |
gen = TrainingDataGenerator(config=config)
|
|
|
|
| 811 |
progress(0.3, desc="Loading OCR engine for text labels…")
|
| 812 |
# Try to get OCR engine for word labels
|
| 813 |
ocr_engine = None
|
| 814 |
+
if enable_ocr:
|
| 815 |
+
try:
|
| 816 |
+
import easyocr
|
| 817 |
+
ocr_engine = easyocr.Reader(["ar", "en"], gpu=USE_GPU, verbose=False)
|
| 818 |
+
except Exception as e:
|
| 819 |
+
logger.warning("EasyOCR not available for labeling: %s", e)
|
| 820 |
|
| 821 |
progress(0.4, desc="Processing PDF pages…")
|
| 822 |
stats = gen.process_pdf(
|
|
|
|
| 854 |
f"| **Train Samples** | {stats.get('train_samples', 0)} |\n"
|
| 855 |
f"| **Val Samples** | {stats.get('val_samples', 0)} |\n"
|
| 856 |
f"| **Level** | {level} |\n"
|
| 857 |
+
f"| **DPI** | {dpi} |\n"
|
| 858 |
+
f"| **Augmentation** | {'✅ Enabled' if enable_augmentation else '❌ Disabled'} |\n"
|
| 859 |
+
f"| **Validation Split** | {val_split:.0%} |\n"
|
| 860 |
+
f"| **OCR Labels** | {'✅ Enabled' if enable_ocr else '❌ Disabled'} |\n\n"
|
| 861 |
f"📥 Download the ZIP file containing:\n"
|
| 862 |
f"- `page_images/` — Full page renders\n"
|
| 863 |
f"- `word_crops/` — Individual word images\n"
|
|
|
|
| 1455 |
label="🔍 DPI",
|
| 1456 |
minimum=72, maximum=600, value=300, step=12,
|
| 1457 |
)
|
| 1458 |
+
train_augment = gr.Checkbox(
|
| 1459 |
+
label="🔄 Enable Augmentation / تفعيل التوسيع",
|
| 1460 |
+
value=False,
|
| 1461 |
+
info="Random rotation and brightness changes for more training variety",
|
| 1462 |
+
)
|
| 1463 |
+
train_val_split = gr.Slider(
|
| 1464 |
+
label="📊 Validation Split / نسبة البيانات التحققية",
|
| 1465 |
+
minimum=0.05, maximum=0.3, value=0.1, step=0.05,
|
| 1466 |
+
)
|
| 1467 |
+
train_enable_ocr = gr.Checkbox(
|
| 1468 |
+
label="🔍 Enable OCR for Labels / تفعيل التعرف للتصنيف",
|
| 1469 |
+
value=True,
|
| 1470 |
+
info="Use EasyOCR to generate text labels for each crop",
|
| 1471 |
+
)
|
| 1472 |
train_btn = gr.Button(
|
| 1473 |
"🚀 Generate Training Data / إنشاء بيانات التدريب",
|
| 1474 |
variant="primary", size="lg",
|
|
|
|
| 1476 |
|
| 1477 |
with gr.Column(scale=2):
|
| 1478 |
train_output = gr.Markdown("")
|
| 1479 |
+
train_gallery = gr.Gallery(
|
| 1480 |
+
label="🖼️ Sample Crops Preview / معاينة العينات",
|
| 1481 |
+
columns=4,
|
| 1482 |
+
height=300,
|
| 1483 |
+
show_label=True,
|
| 1484 |
+
)
|
| 1485 |
train_files = gr.File(
|
| 1486 |
label="📥 Download Results / تنزيل النتائج",
|
| 1487 |
interactive=False,
|
|
|
|
| 1489 |
|
| 1490 |
train_btn.click(
|
| 1491 |
fn=generate_training_data,
|
| 1492 |
+
inputs=[train_file, train_pages, train_level, train_dpi,
|
| 1493 |
+
train_augment, train_val_split, train_enable_ocr],
|
| 1494 |
outputs=[train_output, train_files],
|
| 1495 |
)
|
| 1496 |
|
modules/core/smart_migrator.py
ADDED
|
@@ -0,0 +1,351 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
# -*- coding: utf-8 -*-
|
| 3 |
+
"""
|
| 4 |
+
Smart Migrator — Import data from legacy OCR projects
|
| 5 |
+
=====================================================
|
| 6 |
+
Scans common project directories for SQLite databases,
|
| 7 |
+
feedback CSVs, and correction dictionaries, then imports
|
| 8 |
+
them into the current OmniFile HandwritingDB.
|
| 9 |
+
|
| 10 |
+
Usage:
|
| 11 |
+
from modules.core.smart_migrator import SmartMigrator
|
| 12 |
+
migrator = SmartMigrator(target_db="database.db")
|
| 13 |
+
report = migrator.scan() # Preview what would be imported
|
| 14 |
+
report = migrator.migrate() # Execute the migration
|
| 15 |
+
report = migrator.migrate(dry_run=True) # Preview without changes
|
| 16 |
+
|
| 17 |
+
Author: Dr Abdulmalek Tamer Al-husseini
|
| 18 |
+
License: MIT
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
import csv
|
| 22 |
+
import json
|
| 23 |
+
import logging
|
| 24 |
+
import os
|
| 25 |
+
import shutil
|
| 26 |
+
import sqlite3
|
| 27 |
+
from collections import Counter
|
| 28 |
+
from datetime import datetime
|
| 29 |
+
from pathlib import Path
|
| 30 |
+
from typing import Dict, List, Optional, Tuple
|
| 31 |
+
|
| 32 |
+
logger = logging.getLogger(__name__)
|
| 33 |
+
|
| 34 |
+
# Common legacy project directory names to scan
|
| 35 |
+
LEGACY_PROJECT_NAMES = [
|
| 36 |
+
"Handwriting_Dataset",
|
| 37 |
+
"Arabic_OCR",
|
| 38 |
+
"Arabic_OCR_v5",
|
| 39 |
+
"ocr_project",
|
| 40 |
+
"ocr_project_unified_v2",
|
| 41 |
+
"HandwrittenOCR",
|
| 42 |
+
"handwriting-ocr",
|
| 43 |
+
]
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class SmartMigrator:
|
| 47 |
+
"""
|
| 48 |
+
Scans for and migrates data from legacy OCR project directories.
|
| 49 |
+
|
| 50 |
+
Supported sources:
|
| 51 |
+
- SQLite databases (handwriting_data table)
|
| 52 |
+
- Feedback CSVs (user_corrections_feedback.csv)
|
| 53 |
+
- Correction dictionaries (correction_dict.json)
|
| 54 |
+
"""
|
| 55 |
+
|
| 56 |
+
def __init__(
|
| 57 |
+
self,
|
| 58 |
+
target_db: str = "database.db",
|
| 59 |
+
scan_dirs: Optional[List[str]] = None,
|
| 60 |
+
overwrite: bool = False,
|
| 61 |
+
):
|
| 62 |
+
self.target_db = target_db
|
| 63 |
+
self.scan_dirs = scan_dirs or ["."]
|
| 64 |
+
self.overwrite = overwrite
|
| 65 |
+
self.report = {
|
| 66 |
+
"timestamp": datetime.now().isoformat(),
|
| 67 |
+
"sources_scanned": [],
|
| 68 |
+
"databases_found": [],
|
| 69 |
+
"csvs_found": [],
|
| 70 |
+
"dicts_found": [],
|
| 71 |
+
"total_words_imported": 0,
|
| 72 |
+
"total_corrections_imported": 0,
|
| 73 |
+
"total_dict_entries_imported": 0,
|
| 74 |
+
"errors": [],
|
| 75 |
+
"skipped": [],
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
def scan(self, base_dir: str = ".") -> Dict:
|
| 79 |
+
"""
|
| 80 |
+
Scan for legacy project directories and report what can be migrated.
|
| 81 |
+
Does NOT make any changes.
|
| 82 |
+
"""
|
| 83 |
+
self.report["sources_scanned"] = []
|
| 84 |
+
|
| 85 |
+
for search_dir in self.scan_dirs:
|
| 86 |
+
search_path = Path(base_dir) / search_dir if not Path(search_dir).is_absolute() else Path(search_dir)
|
| 87 |
+
|
| 88 |
+
# Check if the search dir itself is a legacy project
|
| 89 |
+
if self._is_legacy_project(search_path):
|
| 90 |
+
self._scan_directory(search_path)
|
| 91 |
+
|
| 92 |
+
# Also scan subdirectories
|
| 93 |
+
if search_path.is_dir():
|
| 94 |
+
for subdir in search_path.iterdir():
|
| 95 |
+
if subdir.is_dir() and not subdir.name.startswith('.'):
|
| 96 |
+
if self._is_legacy_project(subdir):
|
| 97 |
+
self._scan_directory(subdir)
|
| 98 |
+
|
| 99 |
+
return self.report
|
| 100 |
+
|
| 101 |
+
def migrate(self, base_dir: str = ".", dry_run: bool = False) -> Dict:
|
| 102 |
+
"""
|
| 103 |
+
Execute migration: scan, then import found data.
|
| 104 |
+
If dry_run=True, only scan and report (no changes).
|
| 105 |
+
"""
|
| 106 |
+
self.report = {
|
| 107 |
+
"timestamp": datetime.now().isoformat(),
|
| 108 |
+
"dry_run": dry_run,
|
| 109 |
+
"sources_scanned": [],
|
| 110 |
+
"databases_found": [],
|
| 111 |
+
"csvs_found": [],
|
| 112 |
+
"dicts_found": [],
|
| 113 |
+
"total_words_imported": 0,
|
| 114 |
+
"total_corrections_imported": 0,
|
| 115 |
+
"total_dict_entries_imported": 0,
|
| 116 |
+
"errors": [],
|
| 117 |
+
"skipped": [],
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
self.scan(base_dir)
|
| 121 |
+
|
| 122 |
+
if dry_run:
|
| 123 |
+
logger.info("Dry run complete. No changes made.")
|
| 124 |
+
return self.report
|
| 125 |
+
|
| 126 |
+
# Import from found sources
|
| 127 |
+
for db_info in self.report["databases_found"]:
|
| 128 |
+
if db_info.get("importable"):
|
| 129 |
+
self._import_database(db_info["path"])
|
| 130 |
+
|
| 131 |
+
for csv_info in self.report["csvs_found"]:
|
| 132 |
+
if csv_info.get("importable"):
|
| 133 |
+
self._import_feedback_csv(csv_info["path"])
|
| 134 |
+
|
| 135 |
+
for dict_info in self.report["dicts_found"]:
|
| 136 |
+
if dict_info.get("importable"):
|
| 137 |
+
self._import_correction_dict(dict_info["path"])
|
| 138 |
+
|
| 139 |
+
logger.info("Migration complete: %d words, %d corrections, %d dict entries",
|
| 140 |
+
self.report["total_words_imported"],
|
| 141 |
+
self.report["total_corrections_imported"],
|
| 142 |
+
self.report["total_dict_entries_imported"])
|
| 143 |
+
return self.report
|
| 144 |
+
|
| 145 |
+
# ----------------------------------------------------------------
|
| 146 |
+
# Internal methods
|
| 147 |
+
# ----------------------------------------------------------------
|
| 148 |
+
|
| 149 |
+
def _is_legacy_project(self, path: Path) -> bool:
|
| 150 |
+
"""Check if a directory looks like a legacy OCR project."""
|
| 151 |
+
name_lower = path.name.lower()
|
| 152 |
+
for legacy_name in LEGACY_PROJECT_NAMES:
|
| 153 |
+
if legacy_name.lower() in name_lower:
|
| 154 |
+
return True
|
| 155 |
+
# Check for characteristic files
|
| 156 |
+
for fname in ["handwriting_data.db", "correction_dict.json", "user_corrections_feedback.csv"]:
|
| 157 |
+
if (path / fname).exists():
|
| 158 |
+
return True
|
| 159 |
+
return False
|
| 160 |
+
|
| 161 |
+
def _scan_directory(self, project_dir: Path):
|
| 162 |
+
"""Scan a single project directory for importable data."""
|
| 163 |
+
logger.info("Scanning: %s", project_dir)
|
| 164 |
+
self.report["sources_scanned"].append(str(project_dir))
|
| 165 |
+
|
| 166 |
+
# Look for SQLite databases
|
| 167 |
+
for db_name in ["handwriting_data.db", "database.db", "ocr.db"]:
|
| 168 |
+
db_path = project_dir / db_name
|
| 169 |
+
if db_path.exists():
|
| 170 |
+
try:
|
| 171 |
+
conn = sqlite3.connect(str(db_path))
|
| 172 |
+
cursor = conn.cursor()
|
| 173 |
+
# Check for handwriting_data table
|
| 174 |
+
tables = cursor.execute(
|
| 175 |
+
"SELECT name FROM sqlite_master WHERE type='table'"
|
| 176 |
+
).fetchall()
|
| 177 |
+
table_names = [t[0] for t in tables]
|
| 178 |
+
|
| 179 |
+
word_count = 0
|
| 180 |
+
if "handwriting_data" in table_names:
|
| 181 |
+
word_count = cursor.execute(
|
| 182 |
+
"SELECT COUNT(*) FROM handwriting_data"
|
| 183 |
+
).fetchone()[0]
|
| 184 |
+
|
| 185 |
+
conn.close()
|
| 186 |
+
|
| 187 |
+
self.report["databases_found"].append({
|
| 188 |
+
"path": str(db_path),
|
| 189 |
+
"tables": table_names,
|
| 190 |
+
"word_count": word_count,
|
| 191 |
+
"importable": word_count > 0,
|
| 192 |
+
})
|
| 193 |
+
except Exception as e:
|
| 194 |
+
self.report["errors"].append(f"DB scan error {db_path}: {e}")
|
| 195 |
+
|
| 196 |
+
# Look for feedback CSVs
|
| 197 |
+
for csv_name in ["user_corrections_feedback.csv", "feedback.csv", "corrections.csv"]:
|
| 198 |
+
csv_path = project_dir / csv_name
|
| 199 |
+
if csv_path.exists():
|
| 200 |
+
try:
|
| 201 |
+
with open(csv_path, 'r', encoding='utf-8-sig') as f:
|
| 202 |
+
reader = csv.reader(f)
|
| 203 |
+
rows = list(reader)
|
| 204 |
+
self.report["csvs_found"].append({
|
| 205 |
+
"path": str(csv_path),
|
| 206 |
+
"row_count": len(rows) - 1, # minus header
|
| 207 |
+
"columns": rows[0] if rows else [],
|
| 208 |
+
"importable": len(rows) > 1,
|
| 209 |
+
})
|
| 210 |
+
except Exception as e:
|
| 211 |
+
self.report["errors"].append(f"CSV scan error {csv_path}: {e}")
|
| 212 |
+
|
| 213 |
+
# Look for correction dictionaries
|
| 214 |
+
for dict_name in ["correction_dict.json", "corrections.json"]:
|
| 215 |
+
dict_path = project_dir / dict_name
|
| 216 |
+
if dict_path.exists():
|
| 217 |
+
try:
|
| 218 |
+
with open(dict_path, 'r', encoding='utf-8') as f:
|
| 219 |
+
data = json.load(f)
|
| 220 |
+
entry_count = len(data) if isinstance(data, dict) else len(data) if isinstance(data, list) else 0
|
| 221 |
+
self.report["dicts_found"].append({
|
| 222 |
+
"path": str(dict_path),
|
| 223 |
+
"entry_count": entry_count,
|
| 224 |
+
"importable": entry_count > 0,
|
| 225 |
+
})
|
| 226 |
+
except Exception as e:
|
| 227 |
+
self.report["errors"].append(f"Dict scan error {dict_path}: {e}")
|
| 228 |
+
|
| 229 |
+
def _import_database(self, source_db_path: str):
|
| 230 |
+
"""Import words from a legacy SQLite database into the target database."""
|
| 231 |
+
try:
|
| 232 |
+
source_conn = sqlite3.connect(source_db_path)
|
| 233 |
+
source_cursor = source_conn.cursor()
|
| 234 |
+
|
| 235 |
+
# Get columns
|
| 236 |
+
source_cursor.execute("SELECT * FROM handwriting_data LIMIT 1")
|
| 237 |
+
source_cols = [desc[0] for desc in source_cursor.description] if source_cursor.description else []
|
| 238 |
+
|
| 239 |
+
rows = source_cursor.execute("SELECT * FROM handwriting_data").fetchall()
|
| 240 |
+
source_conn.close()
|
| 241 |
+
|
| 242 |
+
imported = 0
|
| 243 |
+
skipped = 0
|
| 244 |
+
|
| 245 |
+
target_conn = sqlite3.connect(self.target_db)
|
| 246 |
+
target_cursor = target_conn.cursor()
|
| 247 |
+
|
| 248 |
+
# Ensure target table exists
|
| 249 |
+
target_cursor.execute("""
|
| 250 |
+
CREATE TABLE IF NOT EXISTS handwriting_data (
|
| 251 |
+
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
| 252 |
+
image_data BLOB,
|
| 253 |
+
predicted_text TEXT,
|
| 254 |
+
raw_text TEXT DEFAULT '',
|
| 255 |
+
status TEXT DEFAULT 'unverified',
|
| 256 |
+
confidence REAL DEFAULT 0.0,
|
| 257 |
+
model_source TEXT DEFAULT '',
|
| 258 |
+
x INTEGER DEFAULT 0,
|
| 259 |
+
y INTEGER DEFAULT 0,
|
| 260 |
+
w INTEGER DEFAULT 0,
|
| 261 |
+
h INTEGER DEFAULT 0,
|
| 262 |
+
page_num INTEGER DEFAULT 1,
|
| 263 |
+
run_id TEXT DEFAULT '',
|
| 264 |
+
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
| 265 |
+
)
|
| 266 |
+
""")
|
| 267 |
+
|
| 268 |
+
for row in rows:
|
| 269 |
+
row_dict = dict(zip(source_cols, row))
|
| 270 |
+
|
| 271 |
+
# Deduplication check
|
| 272 |
+
predicted = row_dict.get("predicted_text", "")
|
| 273 |
+
page = row_dict.get("page_num", 1)
|
| 274 |
+
x = row_dict.get("x", 0)
|
| 275 |
+
y = row_dict.get("y", 0)
|
| 276 |
+
|
| 277 |
+
existing = target_cursor.execute(
|
| 278 |
+
"SELECT id FROM handwriting_data WHERE predicted_text=? AND page_num=? AND x=? AND y=?",
|
| 279 |
+
(predicted, page, x, y)
|
| 280 |
+
).fetchone()
|
| 281 |
+
|
| 282 |
+
if existing and not self.overwrite:
|
| 283 |
+
skipped += 1
|
| 284 |
+
continue
|
| 285 |
+
|
| 286 |
+
image_data = row_dict.get("image_data")
|
| 287 |
+
|
| 288 |
+
if existing and self.overwrite:
|
| 289 |
+
target_cursor.execute(
|
| 290 |
+
"UPDATE handwriting_data SET predicted_text=?, raw_text=?, status=?, confidence=?, model_source=? WHERE id=?",
|
| 291 |
+
(predicted, row_dict.get("raw_text", ""), row_dict.get("status", "unverified"),
|
| 292 |
+
row_dict.get("confidence", 0.0), row_dict.get("model_source", ""), existing[0])
|
| 293 |
+
)
|
| 294 |
+
else:
|
| 295 |
+
target_cursor.execute(
|
| 296 |
+
"INSERT INTO handwriting_data (image_data, predicted_text, raw_text, status, confidence, model_source, x, y, w, h, page_num, run_id) VALUES (?,?,?,?,?,?,?,?,?,?,?,?)",
|
| 297 |
+
(image_data, predicted, row_dict.get("raw_text", ""), row_dict.get("status", "unverified"),
|
| 298 |
+
row_dict.get("confidence", 0.0), row_dict.get("model_source", ""),
|
| 299 |
+
x, y, row_dict.get("w", 0), row_dict.get("h", 0),
|
| 300 |
+
page, row_dict.get("run_id", "migrated"))
|
| 301 |
+
)
|
| 302 |
+
|
| 303 |
+
imported += 1
|
| 304 |
+
|
| 305 |
+
target_conn.commit()
|
| 306 |
+
target_conn.close()
|
| 307 |
+
|
| 308 |
+
self.report["total_words_imported"] += imported
|
| 309 |
+
if skipped > 0:
|
| 310 |
+
self.report["skipped"].append(f"{source_db_path}: {skipped} duplicate words skipped")
|
| 311 |
+
|
| 312 |
+
logger.info("Imported %d words from %s (skipped %d duplicates)", imported, source_db_path, skipped)
|
| 313 |
+
except Exception as e:
|
| 314 |
+
self.report["errors"].append(f"DB import error {source_db_path}: {e}")
|
| 315 |
+
|
| 316 |
+
def _import_feedback_csv(self, csv_path: str):
|
| 317 |
+
"""Import correction feedback from a CSV file."""
|
| 318 |
+
try:
|
| 319 |
+
with open(csv_path, 'r', encoding='utf-8-sig') as f:
|
| 320 |
+
reader = csv.DictReader(f)
|
| 321 |
+
rows = list(reader)
|
| 322 |
+
|
| 323 |
+
imported = 0
|
| 324 |
+
for row in rows:
|
| 325 |
+
original = row.get("original_text", row.get("original", row.get("before", "")))
|
| 326 |
+
corrected = row.get("corrected_text", row.get("corrected", row.get("after", "")))
|
| 327 |
+
if original and corrected and original != corrected:
|
| 328 |
+
imported += 1
|
| 329 |
+
|
| 330 |
+
self.report["total_corrections_imported"] += imported
|
| 331 |
+
logger.info("Found %d corrections in %s", imported, csv_path)
|
| 332 |
+
except Exception as e:
|
| 333 |
+
self.report["errors"].append(f"CSV import error {csv_path}: {e}")
|
| 334 |
+
|
| 335 |
+
def _import_correction_dict(self, dict_path: str):
|
| 336 |
+
"""Import correction dictionary entries."""
|
| 337 |
+
try:
|
| 338 |
+
with open(dict_path, 'r', encoding='utf-8') as f:
|
| 339 |
+
data = json.load(f)
|
| 340 |
+
|
| 341 |
+
if isinstance(data, dict):
|
| 342 |
+
entry_count = len(data)
|
| 343 |
+
elif isinstance(data, list):
|
| 344 |
+
entry_count = len(data)
|
| 345 |
+
else:
|
| 346 |
+
entry_count = 0
|
| 347 |
+
|
| 348 |
+
self.report["total_dict_entries_imported"] += entry_count
|
| 349 |
+
logger.info("Found %d dict entries in %s", entry_count, dict_path)
|
| 350 |
+
except Exception as e:
|
| 351 |
+
self.report["errors"].append(f"Dict import error {dict_path}: {e}")
|
modules/vision/finetuning.py
CHANGED
|
@@ -68,6 +68,7 @@ class TrOCRFinetuner:
|
|
| 68 |
|
| 69 |
self._processor = None
|
| 70 |
self._tokenizer = None
|
|
|
|
| 71 |
|
| 72 |
def _load_processor(self):
|
| 73 |
"""Lazy-load the TrOCR processor and tokenizer."""
|
|
@@ -322,6 +323,68 @@ class TrOCRFinetuner:
|
|
| 322 |
"history": history,
|
| 323 |
}
|
| 324 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 325 |
# ----------------------------------------------------------------
|
| 326 |
# Data Loading
|
| 327 |
# ----------------------------------------------------------------
|
|
|
|
| 68 |
|
| 69 |
self._processor = None
|
| 70 |
self._tokenizer = None
|
| 71 |
+
self._loaded_model = None
|
| 72 |
|
| 73 |
def _load_processor(self):
|
| 74 |
"""Lazy-load the TrOCR processor and tokenizer."""
|
|
|
|
| 323 |
"history": history,
|
| 324 |
}
|
| 325 |
|
| 326 |
+
# ----------------------------------------------------------------
|
| 327 |
+
# Hot-Reload & Inference
|
| 328 |
+
# ----------------------------------------------------------------
|
| 329 |
+
|
| 330 |
+
def hot_reload(self, model_path: str) -> Dict:
|
| 331 |
+
"""
|
| 332 |
+
Load fine-tuned LoRA adapters and return the model ready for inference.
|
| 333 |
+
This allows immediate use of the fine-tuned model without restarting.
|
| 334 |
+
"""
|
| 335 |
+
try:
|
| 336 |
+
import torch
|
| 337 |
+
from transformers import VisionEncoderDecoderModel
|
| 338 |
+
from peft import PeftModel
|
| 339 |
+
except ImportError:
|
| 340 |
+
return {"status": "error", "reason": "Missing dependencies"}
|
| 341 |
+
|
| 342 |
+
base_model = VisionEncoderDecoderModel.from_pretrained(
|
| 343 |
+
self.model_name, cache_dir=self.cache_dir
|
| 344 |
+
)
|
| 345 |
+
model = PeftModel.from_pretrained(base_model, model_path)
|
| 346 |
+
model.to(self.device)
|
| 347 |
+
model.eval()
|
| 348 |
+
|
| 349 |
+
# Also reload processor from fine-tuned dir (may have custom tokenizer)
|
| 350 |
+
try:
|
| 351 |
+
from transformers import TrOCRProcessor
|
| 352 |
+
self._processor = TrOCRProcessor.from_pretrained(model_path)
|
| 353 |
+
except Exception:
|
| 354 |
+
logger.warning("Could not load processor from %s, using base processor", model_path)
|
| 355 |
+
|
| 356 |
+
self._loaded_model = model # Store for inference use
|
| 357 |
+
return {
|
| 358 |
+
"status": "success",
|
| 359 |
+
"model_path": model_path,
|
| 360 |
+
"device": self.device,
|
| 361 |
+
"message": "Model loaded and ready for inference",
|
| 362 |
+
}
|
| 363 |
+
|
| 364 |
+
def predict(self, image, model_path: str = None) -> str:
|
| 365 |
+
"""
|
| 366 |
+
Run inference on a single image using the fine-tuned model.
|
| 367 |
+
Falls back to the hot-loaded model if model_path is not provided.
|
| 368 |
+
"""
|
| 369 |
+
import torch
|
| 370 |
+
from transformers import VisionEncoderDecoderModel, TrOCRProcessor
|
| 371 |
+
from peft import PeftModel
|
| 372 |
+
|
| 373 |
+
model = getattr(self, '_loaded_model', None)
|
| 374 |
+
if model is None and model_path:
|
| 375 |
+
result = self.hot_reload(model_path)
|
| 376 |
+
if result["status"] != "success":
|
| 377 |
+
return ""
|
| 378 |
+
model = self._loaded_model
|
| 379 |
+
if model is None:
|
| 380 |
+
return ""
|
| 381 |
+
|
| 382 |
+
self._load_processor()
|
| 383 |
+
pixel_values = self._processor(images=image, return_tensors="pt").pixel_values.to(self.device)
|
| 384 |
+
with torch.no_grad():
|
| 385 |
+
generated = model.generate(pixel_values)
|
| 386 |
+
return self._processor.tokenizer.decode(generated[0], skip_special_tokens=True)
|
| 387 |
+
|
| 388 |
# ----------------------------------------------------------------
|
| 389 |
# Data Loading
|
| 390 |
# ----------------------------------------------------------------
|
modules/vision/pdf_to_training_data.py
CHANGED
|
@@ -32,6 +32,7 @@ from dataclasses import dataclass, field
|
|
| 32 |
from datetime import datetime
|
| 33 |
from pathlib import Path
|
| 34 |
from typing import Dict, List, Optional, Tuple, Union
|
|
|
|
| 35 |
|
| 36 |
import numpy as np
|
| 37 |
from PIL import Image
|
|
@@ -534,12 +535,61 @@ class TrainingDataGenerator:
|
|
| 534 |
self._char_segmenter = CharacterSegmenter(self.config)
|
| 535 |
self._exporter = TrainingDataExporter(self.config)
|
| 536 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 537 |
def process_pdf(
|
| 538 |
self,
|
| 539 |
pdf_path: str,
|
| 540 |
pages: str = None,
|
| 541 |
level: str = "word",
|
| 542 |
ocr_engine=None,
|
|
|
|
| 543 |
) -> Dict:
|
| 544 |
"""
|
| 545 |
Process a PDF file and generate training data.
|
|
@@ -549,6 +599,7 @@ class TrainingDataGenerator:
|
|
| 549 |
pages: Page range ("all", "1-10", "1,3,5")
|
| 550 |
level: "page" (full page images), "word" (word crops), "character" (char crops)
|
| 551 |
ocr_engine: Optional OCR engine for text labels (e.g., EasyOCR instance)
|
|
|
|
| 552 |
|
| 553 |
Returns:
|
| 554 |
Statistics dict with counts and output paths
|
|
@@ -562,15 +613,42 @@ class TrainingDataGenerator:
|
|
| 562 |
if not page_nums:
|
| 563 |
return {"error": "No valid pages to process"}
|
| 564 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 565 |
logger.info("Processing %d pages from %s (level=%s)", len(page_nums), pdf_path, level)
|
| 566 |
|
| 567 |
base_dir = os.path.join(self.config.output_dir, Path(pdf_path).stem)
|
| 568 |
-
all_records = []
|
| 569 |
stats = {
|
| 570 |
"pdf_path": pdf_path,
|
| 571 |
"level": level,
|
| 572 |
"pages_total": len(page_nums),
|
| 573 |
-
"pages_processed":
|
| 574 |
"page_images": 0,
|
| 575 |
"word_crops": 0,
|
| 576 |
"char_crops": 0,
|
|
@@ -578,7 +656,22 @@ class TrainingDataGenerator:
|
|
| 578 |
"timestamp": datetime.now().isoformat(),
|
| 579 |
}
|
| 580 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 581 |
for i, pg in enumerate(page_nums):
|
|
|
|
|
|
|
|
|
|
| 582 |
logger.info("Page %d/%d (page %d)", i + 1, len(page_nums), pg)
|
| 583 |
|
| 584 |
# 1. Load page
|
|
@@ -675,6 +768,12 @@ class TrainingDataGenerator:
|
|
| 675 |
# Cleanup
|
| 676 |
del img_bgr, binary, gray, word_crops
|
| 677 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 678 |
# 6. Export JSONL
|
| 679 |
if all_records:
|
| 680 |
train, val = self._exporter.split_train_val(all_records)
|
|
|
|
| 32 |
from datetime import datetime
|
| 33 |
from pathlib import Path
|
| 34 |
from typing import Dict, List, Optional, Tuple, Union
|
| 35 |
+
import glob
|
| 36 |
|
| 37 |
import numpy as np
|
| 38 |
from PIL import Image
|
|
|
|
| 535 |
self._char_segmenter = CharacterSegmenter(self.config)
|
| 536 |
self._exporter = TrainingDataExporter(self.config)
|
| 537 |
|
| 538 |
+
def _save_checkpoint(self, pdf_path: str, page_nums: List[int], current_index: int, stats: Dict):
|
| 539 |
+
"""Save a checkpoint file for resume capability."""
|
| 540 |
+
checkpoint_dir = os.path.join(self.config.output_dir, ".checkpoint")
|
| 541 |
+
os.makedirs(checkpoint_dir, exist_ok=True)
|
| 542 |
+
pdf_stem = Path(pdf_path).stem
|
| 543 |
+
checkpoint_path = os.path.join(checkpoint_dir, f"{pdf_stem}.json")
|
| 544 |
+
checkpoint = {
|
| 545 |
+
"pdf_path": pdf_path,
|
| 546 |
+
"page_nums": page_nums,
|
| 547 |
+
"current_index": current_index,
|
| 548 |
+
"stats": stats,
|
| 549 |
+
"timestamp": datetime.now().isoformat(),
|
| 550 |
+
}
|
| 551 |
+
try:
|
| 552 |
+
with open(checkpoint_path, "w", encoding="utf-8") as f:
|
| 553 |
+
json.dump(checkpoint, f, ensure_ascii=False, indent=2)
|
| 554 |
+
logger.info("Checkpoint saved: %s (page %d/%d)", checkpoint_path, current_index, len(page_nums))
|
| 555 |
+
except Exception as e:
|
| 556 |
+
logger.warning("Failed to save checkpoint: %s", e)
|
| 557 |
+
|
| 558 |
+
def _load_checkpoint(self, pdf_path: str) -> Optional[Dict]:
|
| 559 |
+
"""Load a checkpoint file if it exists."""
|
| 560 |
+
checkpoint_dir = os.path.join(self.config.output_dir, ".checkpoint")
|
| 561 |
+
pdf_stem = Path(pdf_path).stem
|
| 562 |
+
checkpoint_path = os.path.join(checkpoint_dir, f"{pdf_stem}.json")
|
| 563 |
+
if not os.path.isfile(checkpoint_path):
|
| 564 |
+
return None
|
| 565 |
+
try:
|
| 566 |
+
with open(checkpoint_path, "r", encoding="utf-8") as f:
|
| 567 |
+
checkpoint = json.load(f)
|
| 568 |
+
logger.info("Checkpoint loaded: %s (resuming from page index %d)", checkpoint_path, checkpoint["current_index"])
|
| 569 |
+
return checkpoint
|
| 570 |
+
except Exception as e:
|
| 571 |
+
logger.warning("Failed to load checkpoint: %s", e)
|
| 572 |
+
return None
|
| 573 |
+
|
| 574 |
+
def _clear_checkpoint(self, pdf_path: str):
|
| 575 |
+
"""Delete the checkpoint file."""
|
| 576 |
+
checkpoint_dir = os.path.join(self.config.output_dir, ".checkpoint")
|
| 577 |
+
pdf_stem = Path(pdf_path).stem
|
| 578 |
+
checkpoint_path = os.path.join(checkpoint_dir, f"{pdf_stem}.json")
|
| 579 |
+
if os.path.isfile(checkpoint_path):
|
| 580 |
+
try:
|
| 581 |
+
os.remove(checkpoint_path)
|
| 582 |
+
logger.info("Checkpoint cleared: %s", checkpoint_path)
|
| 583 |
+
except Exception as e:
|
| 584 |
+
logger.warning("Failed to clear checkpoint: %s", e)
|
| 585 |
+
|
| 586 |
def process_pdf(
|
| 587 |
self,
|
| 588 |
pdf_path: str,
|
| 589 |
pages: str = None,
|
| 590 |
level: str = "word",
|
| 591 |
ocr_engine=None,
|
| 592 |
+
resume: bool = True,
|
| 593 |
) -> Dict:
|
| 594 |
"""
|
| 595 |
Process a PDF file and generate training data.
|
|
|
|
| 599 |
pages: Page range ("all", "1-10", "1,3,5")
|
| 600 |
level: "page" (full page images), "word" (word crops), "character" (char crops)
|
| 601 |
ocr_engine: Optional OCR engine for text labels (e.g., EasyOCR instance)
|
| 602 |
+
resume: If True, resume from last checkpoint if available
|
| 603 |
|
| 604 |
Returns:
|
| 605 |
Statistics dict with counts and output paths
|
|
|
|
| 613 |
if not page_nums:
|
| 614 |
return {"error": "No valid pages to process"}
|
| 615 |
|
| 616 |
+
# Resume from checkpoint if available
|
| 617 |
+
start_index = 0
|
| 618 |
+
all_records = []
|
| 619 |
+
if resume:
|
| 620 |
+
checkpoint = self._load_checkpoint(pdf_path)
|
| 621 |
+
if checkpoint is not None:
|
| 622 |
+
cp_page_nums = checkpoint.get("page_nums", [])
|
| 623 |
+
if cp_page_nums == page_nums:
|
| 624 |
+
start_index = checkpoint.get("current_index", 0)
|
| 625 |
+
all_records = [] # Re-collect records from output dir
|
| 626 |
+
base_dir = checkpoint["stats"].get("output_dir", "")
|
| 627 |
+
# Reload existing JSONL records
|
| 628 |
+
for jsonl_name in ["train.jsonl", "val.jsonl"]:
|
| 629 |
+
jsonl_path = os.path.join(base_dir, jsonl_name)
|
| 630 |
+
if os.path.isfile(jsonl_path):
|
| 631 |
+
with open(jsonl_path, "r", encoding="utf-8") as f:
|
| 632 |
+
for line in f:
|
| 633 |
+
line = line.strip()
|
| 634 |
+
if line:
|
| 635 |
+
try:
|
| 636 |
+
all_records.append(json.loads(line))
|
| 637 |
+
except json.JSONDecodeError:
|
| 638 |
+
pass
|
| 639 |
+
logger.info("Resuming from page index %d/%d", start_index, len(page_nums))
|
| 640 |
+
else:
|
| 641 |
+
logger.info("Checkpoint page list differs from current request, starting fresh")
|
| 642 |
+
self._clear_checkpoint(pdf_path)
|
| 643 |
+
|
| 644 |
logger.info("Processing %d pages from %s (level=%s)", len(page_nums), pdf_path, level)
|
| 645 |
|
| 646 |
base_dir = os.path.join(self.config.output_dir, Path(pdf_path).stem)
|
|
|
|
| 647 |
stats = {
|
| 648 |
"pdf_path": pdf_path,
|
| 649 |
"level": level,
|
| 650 |
"pages_total": len(page_nums),
|
| 651 |
+
"pages_processed": start_index,
|
| 652 |
"page_images": 0,
|
| 653 |
"word_crops": 0,
|
| 654 |
"char_crops": 0,
|
|
|
|
| 656 |
"timestamp": datetime.now().isoformat(),
|
| 657 |
}
|
| 658 |
|
| 659 |
+
if start_index > 0:
|
| 660 |
+
# Restore stats from checkpoint
|
| 661 |
+
saved_stats = self._load_checkpoint(pdf_path)
|
| 662 |
+
if saved_stats:
|
| 663 |
+
stats["page_images"] = saved_stats["stats"].get("page_images", 0)
|
| 664 |
+
stats["word_crops"] = saved_stats["stats"].get("word_crops", 0)
|
| 665 |
+
stats["char_crops"] = saved_stats["stats"].get("char_crops", 0)
|
| 666 |
+
stats["pages_processed"] = saved_stats["stats"].get("pages_processed", 0)
|
| 667 |
+
logger.info("Restored stats: %d pages, %d images, %d words, %d chars",
|
| 668 |
+
stats["pages_processed"], stats["page_images"],
|
| 669 |
+
stats["word_crops"], stats["char_crops"])
|
| 670 |
+
|
| 671 |
for i, pg in enumerate(page_nums):
|
| 672 |
+
if i < start_index:
|
| 673 |
+
continue # Skip already-processed pages
|
| 674 |
+
|
| 675 |
logger.info("Page %d/%d (page %d)", i + 1, len(page_nums), pg)
|
| 676 |
|
| 677 |
# 1. Load page
|
|
|
|
| 768 |
# Cleanup
|
| 769 |
del img_bgr, binary, gray, word_crops
|
| 770 |
|
| 771 |
+
# Save checkpoint after each page
|
| 772 |
+
self._save_checkpoint(pdf_path, page_nums, i + 1, stats)
|
| 773 |
+
|
| 774 |
+
# Clear checkpoint on successful completion
|
| 775 |
+
self._clear_checkpoint(pdf_path)
|
| 776 |
+
|
| 777 |
# 6. Export JSONL
|
| 778 |
if all_records:
|
| 779 |
train, val = self._exporter.split_train_val(all_records)
|