Просмотр исходного кода

added SuryaOCR batch mode; some fixes

master
Evgeniy Ierusalimov 2 недель назад
Родитель
Сommit
7c786c28f6
7 измененных файлов: 84 добавлений и 46 удалений
  1. 3
    3
      README.md
  2. 62
    37
      src/image_ocr.py
  3. 1
    1
      src/image_prepare.py
  4. 14
    1
      src/ocr/surya_engine.py
  5. 2
    2
      src/postprocess/vlm_corrector.py
  6. 1
    1
      src/split/slicer.py
  7. 1
    1
      tests/test_post_crop.py

+ 3
- 3
README.md Просмотреть файл

59
 | `--slice-auto` / `-a` | Автоопределение сетки | — |
59
 | `--slice-auto` / `-a` | Автоопределение сетки | — |
60
 | `--output-dir` / `-o` | Каталог для страниц | `.` |
60
 | `--output-dir` / `-o` | Каталог для страниц | `.` |
61
 | `--pre-rotate` / `-r` | Поворот: 90, 180, 270 | — |
61
 | `--pre-rotate` / `-r` | Поворот: 90, 180, 270 | — |
62
-| `--border` / `-b` | Белая рамка (px) | 5 |
62
+| `--border` / `-b` | Белая рамка (px) | 10 |
63
 | `--unwarp` / `-u` | Выпрямить перспективные искажения (Canny → 4-угольный контур → warp) | OFF |
63
 | `--unwarp` / `-u` | Выпрямить перспективные искажения (Canny → 4-угольный контур → warp) | OFF |
64
 | `--denoise` | Убрать шум сканера (Non-Local Means) | OFF |
64
 | `--denoise` | Убрать шум сканера (Non-Local Means) | OFF |
65
 | `--post-crop` / `--no-post-crop` | Обрезка по контенту | ON |
65
 | `--post-crop` / `--no-post-crop` | Обрезка по контенту | ON |
105
 │                                                         (белая рамка /       │
105
 │                                                         (белая рамка /       │
106
 │                                                         отступ при обрезке   │
106
 │                                                         отступ при обрезке   │
107
 │                                                         содержимого, по      │
107
 │                                                         содержимого, по      │
108
-│                                                         умолчанию 5) 
109
-│                                                         [default: 5] 
108
+│                                                         умолчанию 10)
109
+│                                                         [default: 10]
110
 │ --unwarp      -u                                        Выпрямить            │
110
 │ --unwarp      -u                                        Выпрямить            │
111
 │                                                         перспективные        │
111
 │                                                         перспективные        │
112
 │                                                         искажения перед      │
112
 │                                                         искажения перед      │

+ 62
- 37
src/image_ocr.py Просмотреть файл

53
 
53
 
54
 def _get_ocr_engine(name: str, main_lang: str):
54
 def _get_ocr_engine(name: str, main_lang: str):
55
     if name == "surya":
55
     if name == "surya":
56
-        from src.ocr.surya_engine import surya_ocr_image
57
-        return surya_ocr_image, "Surya 2 VLM"
56
+        from src.ocr.surya_engine import surya_ocr_batch
57
+        return surya_ocr_batch, "Surya 2 VLM", True
58
     if name == "external-qwen3":
58
     if name == "external-qwen3":
59
         from src.ocr.external_vlm_engine import external_vlm_ocr_image
59
         from src.ocr.external_vlm_engine import external_vlm_ocr_image
60
-        return external_vlm_ocr_image, "Qwen3.8 Max (external VLM)"
61
-    return functools.partial(ocr_image, main_lang=main_lang), "PP-StructureV3"
60
+        return external_vlm_ocr_image, "Qwen3.8 Max (external VLM)", False
61
+    return functools.partial(ocr_image, main_lang=main_lang), "PP-StructureV3", False
62
 
62
 
63
 
63
 
64
-def _merge_markdown_files(output_dir: Path, md_files: list[Path]) -> None:
64
+def _merge_markdown_files(output_dir: Path, input_name: str, md_files: list[Path]) -> None:
65
     merged_parts: list[str] = []
65
     merged_parts: list[str] = []
66
     for md_path in sorted(md_files):
66
     for md_path in sorted(md_files):
67
         content = md_path.read_text(encoding="utf-8")
67
         content = md_path.read_text(encoding="utf-8")
68
         merged_parts.append(f"<!-- page: {md_path.stem} -->\n\n{content}")
68
         merged_parts.append(f"<!-- page: {md_path.stem} -->\n\n{content}")
69
     if merged_parts:
69
     if merged_parts:
70
         merged = "# Merged document\n\n" + "\n\n---\n\n".join(merged_parts)
70
         merged = "# Merged document\n\n" + "\n\n---\n\n".join(merged_parts)
71
-        merged_path = output_dir / (output_dir.name + MD_EXTENSION)
71
+        merged_path = output_dir / (input_name + MD_EXTENSION)
72
         merged_path.write_text(merged, encoding="utf-8")
72
         merged_path.write_text(merged, encoding="utf-8")
73
         logger.info("Объединённый документ: %s", merged_path)
73
         logger.info("Объединённый документ: %s", merged_path)
74
 
74
 
75
 
75
 
76
+def _log_low_confidence(raw_json: dict) -> None:
77
+    layout = raw_json.get("layout_det_res", {}).get("boxes", [])
78
+    if not layout:
79
+        return
80
+    low = [(b.get("label", "?"), round(b.get("score", 0), 3)) for b in layout if b.get("score", 1) < 0.5]
81
+    if low:
82
+        labels = ", ".join(f"{lbl}({sc})" for lbl, sc in low)
83
+        logger.warning("Низкая уверенность: %d блоков — %s", len(low), labels)
84
+
85
+
86
+def _apply_vlm(blocks, input_path: Path) -> None:
87
+    from src.postprocess.corrector import apply_corrector
88
+    from src.postprocess.vlm_corrector import vlm_correct_block
89
+    corrector = functools.partial(vlm_correct_block, image_path=str(input_path.resolve()))
90
+    apply_corrector(blocks, corrector, "VLM")
91
+
92
+
93
+def _save_result(page, input_path: Path, output_dir: Path) -> None:
94
+    json_path = output_dir / (input_path.stem + ".json")
95
+    if not json_path.exists():
96
+        json_payload = page.raw_json or {
97
+            "blocks": [{"label": b.label, "content": b.content, "bbox": list(b.bbox)}
98
+                       for b in page.blocks]
99
+        }
100
+        json_path.write_text(json.dumps(json_payload, ensure_ascii=False, indent=2), encoding="utf-8")
101
+        logger.debug("JSON сохранён: %s", json_path)
102
+
103
+    md_path = output_dir / (input_path.stem + MD_EXTENSION)
104
+    generate_markdown(page.blocks, md_path)
105
+    logger.debug("Markdown сохранён: %s", md_path)
106
+    logger.info("%s → %s, %s", input_path.name, md_path.name, json_path.name)
107
+
108
+
76
 def _process_single(
109
 def _process_single(
77
     ocr_fn, input_path: Path,
110
     ocr_fn, input_path: Path,
78
     output_dir_resolved: Path, use_vlm: bool,
111
     output_dir_resolved: Path, use_vlm: bool,
79
 ) -> None:
112
 ) -> None:
80
     """Обрабатывает одно изображение: OCR → fixups → корректоры → JSON + Markdown."""
113
     """Обрабатывает одно изображение: OCR → fixups → корректоры → JSON + Markdown."""
81
-    # Этап OCR — самое долгое и хрупкое место
82
     logger.debug("OCR: начало %s", input_path.name)
114
     logger.debug("OCR: начало %s", input_path.name)
83
     t0 = time.time()
115
     t0 = time.time()
84
     page = ocr_fn(str(input_path))
116
     page = ocr_fn(str(input_path))
86
     logger.info("Распознано %d блоков за %.1fs", len(page.blocks), dt)
118
     logger.info("Распознано %d блоков за %.1fs", len(page.blocks), dt)
87
     logger.debug("OCR: завершено %s (блоков: %d)", input_path.name, len(page.blocks))
119
     logger.debug("OCR: завершено %s (блоков: %d)", input_path.name, len(page.blocks))
88
 
120
 
89
-    # Исправление типичных OCR-ошибок (VaR, V^2)
90
     fix_ocr_errors(page.blocks)
121
     fix_ocr_errors(page.blocks)
91
     logger.debug("Fixups: завершено")
122
     logger.debug("Fixups: завершено")
92
 
123
 
93
-    # Опциональная коррекция через VLM
124
+    _log_low_confidence(page.raw_json)
125
+
94
     if use_vlm:
126
     if use_vlm:
95
-        from src.postprocess.corrector import apply_corrector
96
-        logger.debug("Корректоры: инициализация")
97
-        from src.postprocess.vlm_corrector import vlm_correct_block
98
-        corrector = functools.partial(vlm_correct_block, image_path=str(input_path.resolve()))
99
-        apply_corrector(page.blocks, corrector, "VLM")
100
-        logger.debug("VLM: завершено")
101
-
102
-    # Сохранение JSON (всегда) + Markdown (всегда)
103
-    json_path = output_dir_resolved / (input_path.stem + ".json")
104
-    if not json_path.exists():
105
-        # Если raw_json пуст (Surya), строим из ParsedBlock
106
-        json_payload = page.raw_json or {
107
-            "blocks": [{"label": b.label, "content": b.content, "bbox": list(b.bbox)}
108
-                       for b in page.blocks]
109
-        }
110
-        json_path.write_text(json.dumps(json_payload, ensure_ascii=False, indent=2), encoding="utf-8")
111
-        logger.debug("JSON сохранён: %s", json_path)
127
+        _apply_vlm(page.blocks, input_path)
112
 
128
 
113
-    md_path = output_dir_resolved / (input_path.stem + MD_EXTENSION)
114
-    generate_markdown(page.blocks, md_path)
115
-    logger.debug("Markdown сохранён: %s", md_path)
116
-    logger.info("%s → %s, %s", input_path.name, md_path.name, json_path.name)
129
+    _save_result(page, input_path, output_dir_resolved)
117
 
130
 
118
 
131
 
119
 @app.command()
132
 @app.command()
180
     ] = 5,
193
     ] = 5,
181
 ) -> None:
194
 ) -> None:
182
     """OCR → PostProcess → Markdown. Принимает файл или каталог изображений."""
195
     """OCR → PostProcess → Markdown. Принимает файл или каталог изображений."""
183
-    ocr_fn, engine_name = _get_ocr_engine(ocr_engine, main_lang)
196
+    ocr_fn, engine_name, is_batch = _get_ocr_engine(ocr_engine, main_lang)
184
 
197
 
185
     if input.is_dir():
198
     if input.is_dir():
186
         image_paths = sorted(
199
         image_paths = sorted(
214
 
227
 
215
         total_start = time.time()
228
         total_start = time.time()
216
 
229
 
217
-        for i, img_path in enumerate(image_paths):
218
-            logger.info("[%d/%d] %s...", i + 1 + skipped, len(image_paths) + skipped, img_path.name)
219
-            _process_single(ocr_fn, img_path, out, use_vlm)
220
-            if i < len(image_paths) - 1 and pause > 0:
221
-                time.sleep(pause)
230
+        if is_batch:
231
+            paths = [str(p) for p in image_paths]
232
+            logger.info("Batch-обработка %d изображений...", len(paths))
233
+            pages = ocr_fn(paths)
234
+            for img_path, page in zip(image_paths, pages):
235
+                logger.info("%s → %d блоков", img_path.name, len(page.blocks))
236
+                fix_ocr_errors(page.blocks)
237
+                _log_low_confidence(page.raw_json)
238
+                if use_vlm:
239
+                    _apply_vlm(page.blocks, img_path)
240
+                _save_result(page, img_path, out)
241
+        else:
242
+            for i, img_path in enumerate(image_paths):
243
+                logger.info("[%d/%d] %s...", i + 1 + skipped, len(image_paths) + skipped, img_path.name)
244
+                _process_single(ocr_fn, img_path, out, use_vlm)
245
+                if i < len(image_paths) - 1 and pause > 0:
246
+                    time.sleep(pause)
222
 
247
 
223
         total_dt = time.time() - total_start
248
         total_dt = time.time() - total_start
224
         processed = len(image_paths)
249
         processed = len(image_paths)
225
         avg_dt = total_dt / processed if processed else 0
250
         avg_dt = total_dt / processed if processed else 0
226
 
251
 
227
         if result_one_document:
252
         if result_one_document:
228
-            _merge_markdown_files(out, sorted(out.glob("*.md")))
253
+            _merge_markdown_files(out, input.name, sorted(out.glob("*.md")))
229
 
254
 
230
         logger.info(
255
         logger.info(
231
             "Готово: %d изображений, всего %.1fs, среднее %.1fs/изобр",
256
             "Готово: %d изображений, всего %.1fs, среднее %.1fs/изобр",

+ 1
- 1
src/image_prepare.py Просмотреть файл

140
             "--border",
140
             "--border",
141
             "-b",
141
             "-b",
142
             min=0,
142
             min=0,
143
-            help="Отступ в пикселях (белая рамка / отступ при обрезке содержимого, по умолчанию 5)",
143
+            help="Отступ в пикселях (белая рамка / отступ при обрезке содержимого, по умолчанию 10)",
144
         ),
144
         ),
145
     ] = DEFAULT_BORDER_PX,
145
     ] = DEFAULT_BORDER_PX,
146
     unwarp: Annotated[
146
     unwarp: Annotated[

+ 14
- 1
src/ocr/surya_engine.py Просмотреть файл

72
         image = Image.open(image_path)
72
         image = Image.open(image_path)
73
         return self._process_image(image)
73
         return self._process_image(image)
74
 
74
 
75
+    def process_batch(self, image_paths: list[str]) -> list[OcrPageResult]:
76
+        from PIL import Image
77
+
78
+        images = [_pre_scale(Image.open(p)) for p in image_paths]
79
+        results = self._predictor(images)
80
+        return [self._to_result(results[i], images[i]) if i < len(results) else OcrPageResult()
81
+                for i in range(len(images))]
82
+
75
     def _health_check(self) -> bool:
83
     def _health_check(self) -> bool:
76
         import requests
84
         import requests
77
 
85
 
130
 
138
 
131
 
139
 
132
 def surya_ocr_image(image_path: str) -> OcrPageResult:
140
 def surya_ocr_image(image_path: str) -> OcrPageResult:
141
+    engine = _get_surya_engine()
142
+    return engine.process(image_path)
143
+
144
+
145
+def surya_ocr_batch(image_paths: list[str]) -> list[OcrPageResult]:
133
     engine = _get_surya_engine()
146
     engine = _get_surya_engine()
134
     if not engine._health_check():
147
     if not engine._health_check():
135
         logger.warning("llama-server не отвечает, перезапуск...")
148
         logger.warning("llama-server не отвечает, перезапуск...")
141
                 "Проверьте: `ps aux | grep llama-server`. "
154
                 "Проверьте: `ps aux | grep llama-server`. "
142
                 "Попробуйте: `llama-server --version`"
155
                 "Попробуйте: `llama-server --version`"
143
             )
156
             )
144
-    return engine.process(image_path)
157
+    return engine.process_batch(image_paths)

+ 2
- 2
src/postprocess/vlm_corrector.py Просмотреть файл

26
         return
26
         return
27
 
27
 
28
     crop = image[y1:y2, x1:x2]
28
     crop = image[y1:y2, x1:x2]
29
-    logger.info("VLM: отправка региона %dx%d px (bbox %d,%d,%d,%d)", x2 - x1, y2 - y1, x1, y1, x2, y2)
30
-
31
     b64 = image_to_base64(crop)
29
     b64 = image_to_base64(crop)
30
+    total_sent = len(b64) // 1024
31
+    logger.info("VLM: отправка %dx%d px → %d KB (bbox %d,%d,%d,%d)", x2 - x1, y2 - y1, total_sent, x1, y1, x2, y2)
32
     text = call_qwen_vlm(
32
     text = call_qwen_vlm(
33
         b64,
33
         b64,
34
         "Extract ALL visible text from this image region. Preserve LaTeX math. Return ONLY the extracted text.",
34
         "Extract ALL visible text from this image region. Preserve LaTeX math. Return ONLY the extracted text.",

+ 1
- 1
src/split/slicer.py Просмотреть файл

10
 
10
 
11
 SUPPORTED_EXTENSIONS: frozenset[str] = frozenset({".jpg", ".jpeg", ".png", ".tiff", ".tif"})
11
 SUPPORTED_EXTENSIONS: frozenset[str] = frozenset({".jpg", ".jpeg", ".png", ".tiff", ".tif"})
12
 PAGE_NUMBER_WIDTH: int = 2
12
 PAGE_NUMBER_WIDTH: int = 2
13
-DEFAULT_BORDER_PX: int = 5
13
+DEFAULT_BORDER_PX: int = 10
14
 BORDER_COLOR: tuple[int, int, int] = (255, 255, 255)
14
 BORDER_COLOR: tuple[int, int, int] = (255, 255, 255)
15
 VALID_ROTATIONS: frozenset[int] = frozenset({90, 180, 270})
15
 VALID_ROTATIONS: frozenset[int] = frozenset({90, 180, 270})
16
 DENSITY_RATIO: float = 0.005
16
 DENSITY_RATIO: float = 0.005

+ 1
- 1
tests/test_post_crop.py Просмотреть файл

121
             result = process_image(Path(path), 1, 1, Path(tmp), post_crop=False)
121
             result = process_image(Path(path), 1, 1, Path(tmp), post_crop=False)
122
             assert len(result) == 1
122
             assert len(result) == 1
123
             loaded = cv2.imread(str(result[0]))
123
             loaded = cv2.imread(str(result[0]))
124
-            assert loaded.shape == (310, 410, 3)
124
+            assert loaded.shape == (320, 420, 3)
125
 
125
 
126
     def test_post_crop_e2e_2x2(self) -> None:
126
     def test_post_crop_e2e_2x2(self) -> None:
127
         image = create_test_image(400, 400)
127
         image = create_test_image(400, 400)

Загрузка…
Отмена
Сохранить