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 CHANGED
@@ -11,7 +11,7 @@ license: mit
11
 
12
  <div align="center">
13
 
14
- # 🧠 OmniFile AI Processor v4.3.0
15
 
16
  **نظام ذكاء اصطناعي متكامل لمعالجة الملفات والنصوص والخط اليدوي**
17
  **A Comprehensive AI System for File Processing, Text Analysis & Handwriting Recognition**
@@ -23,7 +23,7 @@ license: mit
23
  [![GitHub](https://img.shields.io/badge/GitHub-DrAbdulmalek-181717?logo=github)](https://github.com/DrAbdulmalek/OmniFile_Processor)
24
 
25
  <p>
26
- <b>Version:</b> v4.3.0 &nbsp;|&nbsp;
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
  [![GitHub](https://img.shields.io/badge/GitHub-DrAbdulmalek-181717?logo=github)](https://github.com/DrAbdulmalek/OmniFile_Processor)
24
 
25
  <p>
26
+ <b>Version:</b> v4.4.0 &nbsp;|&nbsp;
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 + zero-shot
15
  6. Evaluation — CER / WER metrics (Levenshtein + jiwer)
16
- 7. About — Project info, links, author
 
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
- try:
809
- import easyocr
810
- ocr_engine = easyocr.Reader(["ar", "en"], gpu=USE_GPU, verbose=False)
811
- except Exception as e:
812
- logger.warning("EasyOCR not available for labeling: %s", e)
 
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\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": 0,
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)