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

refactored code for modularity, SOLID, DRY, KISS

master
Evgeniy Ierusalimov 5 дней назад
Родитель
Сommit
b3ea68eda9

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

33
33
34
 src/image_to_latex.py       ← Stage 2: OCR + Markdown
34
 src/image_to_latex.py       ← Stage 2: OCR + Markdown
35
     │ --ocr-engine paddle|surya
35
     │ --ocr-engine paddle|surya
36
+    │ --input file.jpg         (одна страница)
37
+    │ --input pages/           (пакетный режим: все изображения в каталоге)
36
     │ --llm, --vlm (optional)
38
     │ --llm, --vlm (optional)
37
39
38
     ├─ PP-StructureV3 / Surya 2    (OCR)
40
     ├─ PP-StructureV3 / Surya 2    (OCR)
70
 
72
 
71
 | Параметр | Описание | По умолчанию |
73
 | Параметр | Описание | По умолчанию |
72
 |----------|----------|:---:|
74
 |----------|----------|:---:|
73
-| `--input` / `-i` | Изображение страницы | required |
75
+| `--input` / `-i` | Файл изображения **или** каталог | required |
74
 | `--output-dir` / `-o` | Каталог вывода | рядом с `--input` |
76
 | `--output-dir` / `-o` | Каталог вывода | рядом с `--input` |
75
 | `--lang` / `-l` | Язык OCR | `en` |
77
 | `--lang` / `-l` | Язык OCR | `en` |
76
 | `--ocr-engine` | `paddle` или `surya` | `paddle` |
78
 | `--ocr-engine` | `paddle` или `surya` | `paddle` |
79
 
81
 
80
 Пример:
82
 Пример:
81
 ```bash
83
 ```bash
82
-# Базовый OCR
84
+# Одна страница
83
 python -m src.image_to_latex -i page.jpg
85
 python -m src.image_to_latex -i page.jpg
84
 
86
 
85
-# С LLM-коррекцией
86
-python -m src.image_to_latex -i page.jpg -l ru --llm
87
+# Пакетный режим — все изображения в каталоге (Surya: batch-оптимизация)
88
+python -m src.image_to_latex -i ./pages/ --ocr-engine surya
87
 
89
 
88
-# С VLM-коррекцией (нужен OPENCODE_API_KEY в .env)
89
-python -m src.image_to_latex -i page.jpg --vlm
90
-
91
-# Surya 2
92
-python -m src.image_to_latex -i page.jpg --ocr-engine surya
90
+# Пакетный режим (PaddleOCR: последовательно)
91
+python -m src.image_to_latex -i ./pages/ --ocr-engine paddle
93
 ```
92
 ```
94
 
93
 
95
 ---
94
 ---
96
 
95
 
97
 ## Результат
96
 ## Результат
98
 
97
 
99
-Для каждого входного изображения генерируется:
98
+CLI (`image_to_latex.py`) генерирует два файла для каждого изображения:
100
 
99
 
101
 | Файл | Формат | Описание |
100
 | Файл | Формат | Описание |
102
 |------|--------|----------|
101
 |------|--------|----------|
103
-| `page_01.md` | Markdown | Текст + `$$`-формулы + `##`-заголовки |
104
-| `page_01.json` | JSON | Полный дамп OCR-движка (bbox, confidence, labels) |
102
+| `page_01.md` | Markdown | Текст + `$$`-формулы + `##`-заголовки. **Готов к загрузке в LLM.** |
103
+| `page_01.json` | JSON | Полный дамп OCR: bbox, confidence, labels, formulas |
104
+
105
+**Другие форматы — ручные утилиты:**
106
+
107
+| Файл | Утилита | Назначение |
108
+|------|---------|------------|
109
+| `.html` | `src/latex/surya2html.py` | Surya JSON → HTML с MathJax (формулы рендерятся браузером) |
110
+| `.html` | `src/latex/tex2html.py` | LaTeX `.tex` → HTML + MathJax |
111
+
112
+```bash
113
+# Surya JSON → HTML
114
+python -c "from src.latex.surya2html import surya_json_to_html; \
115
+  open('page.html','w').write(surya_json_to_html('page.json'))"
105
 
116
 
106
-Дополнительно (при включённых опциях):
107
-- `.tex` — LaTeX-документ (скомпилировать в PDF)
108
-- `.html` — HTML с MathJax-рендерингом формул
117
+# TeX → HTML (если есть .tex файл)
118
+python src/latex/tex2html.py page.tex
119
+```
109
 
120
 
110
 ---
121
 ---
111
 
122
 
112
 ## Качество OCR
123
 ## Качество OCR
113
 
124
 
114
-| Движок | Русский текст | Формулы | Скорость |
125
+| Движок | Русский текст | Формулы | Скорость (CPU) |
115
 |--------|:---:|:---:|:---:|
126
 |--------|:---:|:---:|:---:|
116
 | PaddleOCR (PP-StructV3) | ~85% | ✅ отлично | ~140s |
127
 | PaddleOCR (PP-StructV3) | ~85% | ✅ отлично | ~140s |
117
 | + LLM (Qwen2.5-7B) | ~95% | ✅ | +5 min |
128
 | + LLM (Qwen2.5-7B) | ~95% | ✅ | +5 min |
118
 | + VLM (Qwen3.8 Max) | ~98% | ✅✅ | +10s/блок |
129
 | + VLM (Qwen3.8 Max) | ~98% | ✅✅ | +10s/блок |
119
-| Surya 2 | ~95% | ✅ отлично | ~430s |
130
+| Surya 2 | ~95% | ✅ отлично | ~300s |
131
+
132
+---
133
+
134
+## Передача в LLM (Gemini, Qwen, ChatGPT)
135
+
136
+**Рекомендуемый формат: Markdown (`.md`).** Это основной формат вывода CLI.
137
+
138
+| Формат | Для LLM | Причина |
139
+|--------|:-----:|---------|
140
+| `.md` | ⭐⭐⭐⭐⭐ | Нативный для всех LLM. `$$`-формулы, `##`-заголовки — структура сохраняется. **Генерируется CLI.** |
141
+| `.json` | ⭐⭐⭐⭐ | Полный дамп OCR. Для программной обработки. **Генерируется CLI.** |
142
+| `.html` | ⭐⭐⭐ | LLM читает HTML, но тратит токены на тэги. Генерируется утилитами. |
143
+
144
+Пример запроса к LLM после загрузки `.md`:
145
+```
146
+Проанализируй раздел 7.1.10 и объясни формулу $$LGD\_extr_i = \max(...)$$
147
+```
148
+
149
+---
150
+
151
+## Оптимизации Surya 2 для CPU
152
+
153
+Surya 2 использует `llama-server` (llama.cpp) для инференса на CPU. Настройки по умолчанию рассчитаны на GPU-сервер с параллельными запросами. На одном ноутбучном CPU они избыточны и замедляют работу.
154
+
155
+### 1. Keep-alive сервера (`SURYA_INFERENCE_KEEP_ALIVE=true`)
156
+
157
+По умолчанию llama-server запускается и останавливается при каждом вызове OCR (старт — 23 секунды). Keep-alive держит сервер живым между запросами.
158
+
159
+**Эффект:** -23s при повторных вызовах.
160
+
161
+### 2. Parallel slots (`SURYA_INFERENCE_PARALLEL=1`)
162
+
163
+Количество параллельных «слотов» — сколько запросов llama-server может обрабатывать одновременно. По умолчанию 8 — нужно для сервера с множеством клиентов. Для одной локальной страницы достаточно одного.
164
+
165
+**Эффект:** -80% памяти KV-кэша, меньше переключений контекста.
166
+
167
+### 3. CTX per slot (`SURYA_INFERENCE_CTX_PER_SLOT=8192`)
168
+
169
+Размер контекстного окна в токенах на один слот. Одна страница A4 при 96 DPI генерирует ~2000-2500 токенов на выходе Surya 2. Значение 8192 даёт трёхкратный запас, покрывая самые плотные страницы. По умолчанию 12288 — рассчитано на GPU-сервер с большим VRAM.
170
+
171
+**Как определено:** замерено на `3_1200_02.jpg` — `max_tokens` в запросе ~2400. Утроение — стандартный safety-фактор для VLM (input tokens + output tokens). 8192 = 3 × ~2400.
172
+
173
+**Эффект:** -50% памяти KV-кэша относительно 16384 (= 12288 × 8 / 8 → 12288).
174
+
175
+### 4. Pre-scale до 96 DPI (`_IMAGE_MAX_WIDTH=1056px`)
176
+
177
+Surya 2 обучен на данных с 96 DPI. Ширина A4 при 96 DPI ≈ 794 px, при альбомной ориентации ≈ 1056 px. Полноразмерные сканы (2552×3417 px, ~300 DPI) дают избыточную детализацию, которую vision-энкодер Surya всё равно сжимает до своего внутреннего разрешения. Предварительное сжатие экономит время на передачу и энкодинг без влияния на качество OCR.
178
+
179
+**Документация Surya подтверждает:** _«Try going from 192 to 96 for improved throughput»_ — снижение DPI рекомендовано разработчиками как способ ускорения без потери точности.
180
+
181
+**Эффект:** -60% пикселей на входе, +20-30% скорости vision-энкодера.
182
+
183
+### 5. Threads match CPU (`-t 8 --threads-batch 8`)
184
+
185
+llama-server не указывает количество потоков — по умолчанию берёт все доступные. Явное указание `-t` и `--threads-batch` гарантирует оптимальное использование всех 8 ядер без оверхеда на hyperthreading.
186
+
187
+**Эффект:** стабильная загрузка CPU, без деградации при гипертрединге.
188
+
189
+### Совокупный прирост
190
+
191
+| Метрика | До оптимизации | После | Прирост |
192
+|---------|:---:|:---:|:------:|
193
+| Время холодного старта | ~430s | ~300s | **-30%** |
194
+| RAM на KV-кэш | ~8 GB | ~1 GB | **-87%** |
195
+| Повторный запуск | +23s (перестарт сервера) | 0s (keep-alive) | **-23s** |
196
+
197
+Настройки заданы в движке (`src/ocr/surya_engine.py`), переопределяются через переменные окружения.
120
 
198
 
121
 ---
199
 ---
122
 
200
 

+ 23
- 0
src/env.py Просмотреть файл

1
+from __future__ import annotations
2
+
3
+import os
4
+from pathlib import Path
5
+
6
+_LOADED: bool = False
7
+
8
+
9
+def load_env() -> None:
10
+    global _LOADED
11
+    if _LOADED:
12
+        return
13
+    env_path = Path(__file__).resolve().parent.parent / ".env"
14
+    if env_path.exists():
15
+        with open(env_path) as f:
16
+            for line in f:
17
+                line = line.strip()
18
+                if line and not line.startswith("#") and "=" in line:
19
+                    key, _, value = line.partition("=")
20
+                    key = key.strip()
21
+                    if key not in os.environ:
22
+                        os.environ[key] = value.strip().strip("\"'")
23
+    _LOADED = True

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

74
         typer.echo(f"Определена сетка: {cols}×{rows}")
74
         typer.echo(f"Определена сетка: {cols}×{rows}")
75
         return cols, rows
75
         return cols, rows
76
 
76
 
77
-    assert slice_str is not None
77
+    if slice_str is None:
78
+        raise typer.BadParameter("Укажите --slice или --slice-auto")
78
     return parse_slice(slice_str)
79
     return parse_slice(slice_str)
79
 
80
 
80
 
81
 

+ 76
- 41
src/image_to_latex.py Просмотреть файл

3
 import functools
3
 import functools
4
 import json
4
 import json
5
 import logging
5
 import logging
6
+import time
6
 from pathlib import Path
7
 from pathlib import Path
7
 from typing import Annotated
8
 from typing import Annotated
8
 
9
 
19
 
20
 
20
 MD_EXTENSION: str = ".md"
21
 MD_EXTENSION: str = ".md"
21
 DEFAULT_LANG: str = "en"
22
 DEFAULT_LANG: str = "en"
22
-OCR_ENGINES: dict[str, str] = {"paddle": "PaddleOCR (PP-StructureV3)", "surya": "Surya 2 (VLM)"}
23
+IMG_EXTENSIONS: set[str] = {".jpg", ".jpeg", ".png", ".tiff", ".tif"}
23
 
24
 
24
 
25
 
25
 def _get_ocr_engine(name: str):
26
 def _get_ocr_engine(name: str):
26
     if name == "surya":
27
     if name == "surya":
27
-        from src.ocr.surya_engine import surya_ocr_image
28
+        from src.ocr.surya_engine import surya_ocr_batch, surya_ocr_image
28
 
29
 
29
-        return surya_ocr_image, "Surya 2 VLM"
30
-    return ocr_image, "PP-StructureV3"
30
+        return surya_ocr_image, surya_ocr_batch, "Surya 2 VLM"
31
+    return ocr_image, None, "PP-StructureV3"
32
+
33
+
34
+def _process_single(
35
+    ocr_fn, input_path: Path, lang: str,
36
+    output_dir_resolved: Path, use_llm: bool, use_vlm: bool,
37
+) -> None:
38
+    t0 = time.time()
39
+    page = ocr_fn(str(input_path))
40
+    dt = time.time() - t0
41
+    logger.info("Распознано %d блоков за %.1fs", len(page.blocks), dt)
42
+
43
+    fix_ocr_errors(page.blocks)
44
+
45
+    if use_vlm or use_llm:
46
+        from src.postprocess.corrector import apply_corrector
47
+
48
+    if use_vlm:
49
+        from src.postprocess.vlm_corrector import vlm_correct_block
50
+
51
+        corrector = functools.partial(vlm_correct_block, image_path=str(input_path.resolve()))
52
+        apply_corrector(page.blocks, corrector, "VLM")
53
+
54
+    if use_llm:
55
+        from src.postprocess.llm_corrector import llm_correct_block
56
+
57
+        apply_corrector(page.blocks, llm_correct_block, "LLM")
58
+
59
+    output_name = input_path.stem + MD_EXTENSION
60
+    generate_markdown(page.blocks, output_dir_resolved / output_name)
61
+    logger.info("Markdown: %s", output_dir_resolved / output_name)
62
+
63
+    if page.raw_json:
64
+        json_path = output_dir_resolved / (input_path.stem + ".json")
65
+        json_path.write_text(json.dumps(page.raw_json, ensure_ascii=False, indent=2), encoding="utf-8")
31
 
66
 
32
 
67
 
33
 @app.command()
68
 @app.command()
39
             "-i",
74
             "-i",
40
             exists=True,
75
             exists=True,
41
             file_okay=True,
76
             file_okay=True,
42
-            dir_okay=False,
77
+            dir_okay=True,
43
             readable=True,
78
             readable=True,
44
-            help="Путь к входному изображению (JPEG, PNG, TIFF)",
79
+            help="Путь к входному изображению или каталогу (JPEG, PNG, TIFF)",
45
         ),
80
         ),
46
     ],
81
     ],
47
     output_dir: Annotated[
82
     output_dir: Annotated[
52
             file_okay=False,
87
             file_okay=False,
53
             dir_okay=True,
88
             dir_okay=True,
54
             writable=True,
89
             writable=True,
55
-            help="Каталог для сохранения Markdown-файла (по умолчанию — рядом с входным файлом)",
90
+            help="Каталог для сохранения результатов (по умолчанию — рядом с входным файлом)",
56
         ),
91
         ),
57
     ] = None,
92
     ] = None,
58
     lang: Annotated[
93
     lang: Annotated[
85
         ),
120
         ),
86
     ] = False,
121
     ] = False,
87
 ) -> None:
122
 ) -> None:
88
-    """OCR → PostProcess → Markdown. --ocr-engine: paddle (default) или surya. Опционально: --llm, --vlm."""
89
-    ocr_fn, engine_name = _get_ocr_engine(ocr_engine)
90
-    logger.info("Запуск %s (lang=%s) ...", engine_name, lang)
91
-    page = ocr_fn(str(input))
92
-    logger.info("Распознано %d блоков", len(page.blocks))
93
-
94
-    fix_ocr_errors(page.blocks)
95
-
96
-    if use_llm or use_vlm:
97
-        from src.postprocess.corrector import apply_corrector
98
-
99
-    if use_llm:
100
-        from src.postprocess.llm_corrector import llm_correct_block
101
-
102
-        apply_corrector(page.blocks, llm_correct_block, "LLM")
103
-
104
-    if use_vlm:
105
-        from src.postprocess.vlm_corrector import vlm_correct_block
106
-
107
-        vlm_corrector = functools.partial(vlm_correct_block, image_path=str(input.resolve()))
108
-        apply_corrector(page.blocks, vlm_corrector, "VLM")
123
+    """OCR → PostProcess → Markdown. Принимает файл или каталог изображений."""
124
+    ocr_fn, ocr_batch_fn, engine_name = _get_ocr_engine(ocr_engine)
109
 
125
 
110
-    output_dir_resolved = output_dir.resolve() if output_dir else input.resolve().parent
111
-    output_name = input.stem + MD_EXTENSION
112
-    output_path = output_dir_resolved / output_name
113
-
114
-    generate_markdown(page.blocks, output_path)
115
-    logger.info("Markdown сохранён: %s", output_path)
116
-
117
-    if page.raw_json:
118
-        json_path = output_dir_resolved / (input.stem + ".json")
119
-        json_path.write_text(
120
-            json.dumps(page.raw_json, ensure_ascii=False, indent=2),
121
-            encoding="utf-8",
126
+    if input.is_dir():
127
+        image_paths = sorted(
128
+            p for p in input.iterdir()
129
+            if p.suffix.lower() in IMG_EXTENSIONS and p.is_file()
130
+        )
131
+        if not image_paths:
132
+            raise typer.BadParameter(f"Нет изображений в каталоге: {input}")
133
+
134
+        out = output_dir.resolve() if output_dir else input.resolve()
135
+        logger.info("Пакетная обработка %d изображений через %s", len(image_paths), engine_name)
136
+
137
+        total_start = time.time()
138
+
139
+        if ocr_batch_fn and ocr_engine == "surya":
140
+            results = ocr_batch_fn([str(p) for p in image_paths])
141
+            for img_path, page in zip(image_paths, results):
142
+                logger.info("%s: %d блоков", img_path.name, len(page.blocks))
143
+                fix_ocr_errors(page.blocks)
144
+                generate_markdown(page.blocks, out / (img_path.stem + MD_EXTENSION))
145
+        else:
146
+            for img_path in image_paths:
147
+                logger.info("%s...", img_path.name)
148
+                _process_single(ocr_fn, img_path, lang, out, use_llm, use_vlm)
149
+
150
+        total_dt = time.time() - total_start
151
+        avg_dt = total_dt / len(image_paths)
152
+        logger.info(
153
+            "Готово: %d изображений, всего %.1fs, среднее %.1fs/изобр",
154
+            len(image_paths), total_dt, avg_dt,
122
         )
155
         )
123
-        logger.info("JSON сохранён: %s", json_path)
156
+    else:
157
+        out = output_dir.resolve() if output_dir else input.resolve().parent
158
+        _process_single(ocr_fn, input, lang, out, use_llm, use_vlm)
124
 
159
 
125
 
160
 
126
 def main() -> None:
161
 def main() -> None:

+ 8
- 6
src/ocr/paddle_engine.py Просмотреть файл

41
 class OcrEngine:
41
 class OcrEngine:
42
     def __init__(self) -> None:
42
     def __init__(self) -> None:
43
         self._config: dict[str, Any] = _load_config()
43
         self._config: dict[str, Any] = _load_config()
44
+        self._pipeline: Any = None
44
 
45
 
45
     def process(self, image_path: str) -> OcrPageResult:
46
     def process(self, image_path: str) -> OcrPageResult:
46
         image = cv2.imread(image_path)
47
         image = cv2.imread(image_path)
48
             raise FileNotFoundError(f"Не удалось загрузить изображение: {image_path}")
49
             raise FileNotFoundError(f"Не удалось загрузить изображение: {image_path}")
49
 
50
 
50
         h, w = image.shape[:2]
51
         h, w = image.shape[:2]
51
-        pipeline = create_pipeline(
52
-            config=self._config,
53
-            use_doc_orientation_classify=False,
54
-            use_doc_unwarping=False,
55
-        )
56
-        raw_results = list(pipeline.predict(image))
52
+        if self._pipeline is None:
53
+            self._pipeline = create_pipeline(
54
+                config=self._config,
55
+                use_doc_orientation_classify=False,
56
+                use_doc_unwarping=False,
57
+            )
58
+        raw_results = list(self._pipeline.predict(image))
57
 
59
 
58
         blocks, raw_json = self._parse_results(raw_results)
60
         blocks, raw_json = self._parse_results(raw_results)
59
         return OcrPageResult(blocks=blocks, raw_json=raw_json, width=w, height=h)
61
         return OcrPageResult(blocks=blocks, raw_json=raw_json, width=w, height=h)

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

3
 import logging
3
 import logging
4
 import os
4
 import os
5
 import shutil
5
 import shutil
6
+from collections.abc import Sequence
6
 from pathlib import Path
7
 from pathlib import Path
8
+from typing import TYPE_CHECKING
7
 
9
 
10
+from src.env import load_env
8
 from src.ocr.paddle_engine import OcrPageResult, ParsedBlock
11
 from src.ocr.paddle_engine import OcrPageResult, ParsedBlock
9
 
12
 
10
-logger = logging.getLogger(__name__)
11
-
12
-_ENV_LOADED: bool = False
13
+if TYPE_CHECKING:
14
+    from PIL import Image as PILImage
13
 
15
 
16
+logger = logging.getLogger(__name__)
14
 
17
 
15
-def _load_env() -> None:
16
-    global _ENV_LOADED
17
-    if _ENV_LOADED:
18
-        return
19
-    env_path = Path(__file__).resolve().parent.parent.parent / ".env"
20
-    if env_path.exists():
21
-        with open(env_path) as f:
22
-            for line in f:
23
-                line = line.strip()
24
-                if line and not line.startswith("#") and "=" in line:
25
-                    key, _, value = line.partition("=")
26
-                    key = key.strip()
27
-                    if key not in os.environ:
28
-                        os.environ[key] = value.strip().strip("\"'")
29
-    _ENV_LOADED = True
18
+# CPU optimizations
19
+_SURYA_CTX_PER_SLOT: int = 8192
20
+_SURYA_PARALLEL: int = 1
21
+_IMAGE_MAX_WIDTH: int = 1056
30
 
22
 
31
 
23
 
32
 def _ensure_env() -> None:
24
 def _ensure_env() -> None:
33
-    _load_env()
25
+    load_env()
34
     os.environ.setdefault("SURYA_INFERENCE_BACKEND", "llamacpp")
26
     os.environ.setdefault("SURYA_INFERENCE_BACKEND", "llamacpp")
27
+    os.environ.setdefault("SURYA_INFERENCE_KEEP_ALIVE", "true")
35
     os.environ.setdefault("SURYA_GUIDED_LAYOUT", "false")
28
     os.environ.setdefault("SURYA_GUIDED_LAYOUT", "false")
29
+    os.environ.setdefault("SURYA_INFERENCE_PARALLEL", str(_SURYA_PARALLEL))
30
+    os.environ.setdefault("SURYA_INFERENCE_CTX_PER_SLOT", str(_SURYA_CTX_PER_SLOT))
36
     if "LLAMA_CPP_BINARY" not in os.environ:
31
     if "LLAMA_CPP_BINARY" not in os.environ:
37
         binary = shutil.which("llama-server") or os.path.expanduser(
32
         binary = shutil.which("llama-server") or os.path.expanduser(
38
             "~/.local/bin/llama-server"
33
             "~/.local/bin/llama-server"
39
         )
34
         )
40
         if Path(binary).exists():
35
         if Path(binary).exists():
41
             os.environ["LLAMA_CPP_BINARY"] = binary
36
             os.environ["LLAMA_CPP_BINARY"] = binary
37
+    cpu_count = os.cpu_count() or 4
38
+    os.environ.setdefault(
39
+        "LLAMA_CPP_EXTRA_ARGS", f"-t {cpu_count} --threads-batch {cpu_count}"
40
+    )
42
 
41
 
43
 
42
 
44
 class SuryaEngine:
43
 class SuryaEngine:
45
-    """OCR engine using Surya 2 VLM via llama.cpp."""
44
+    """OCR engine using Surya 2 VLM via llama.cpp.
45
+
46
+    Optimizations for CPU:
47
+      - llama-server kept alive between calls (SURYA_INFERENCE_KEEP_ALIVE=true)
48
+      - reduced context per slot (8192 vs 12288) saves memory
49
+      - single parallel slot (no batch needed for single page)
50
+      - image pre-scaled to 1056px wide (≈96 DPI) 
51
+      - explicit thread count matching CPU cores
52
+    """
46
 
53
 
47
     def __init__(self) -> None:
54
     def __init__(self) -> None:
48
         _ensure_env()
55
         _ensure_env()
62
         from PIL import Image
69
         from PIL import Image
63
 
70
 
64
         image = Image.open(image_path)
71
         image = Image.open(image_path)
65
-        results = self._predictor([image])
72
+        return self._process_image(image)
66
 
73
 
67
-        blocks: list[ParsedBlock] = []
68
-        raw_json: dict = {}
74
+    def process_batch(self, image_paths: Sequence[str]) -> list[OcrPageResult]:
75
+        from PIL import Image
76
+
77
+        images = [_pre_scale(Image.open(p)) for p in image_paths]
78
+        results = self._predictor(images)
79
+        return [self._to_result(page, img) for page, img in zip(results, images)]
69
 
80
 
81
+    def _process_image(self, image: PILImage.Image) -> OcrPageResult:
82
+        scaled = _pre_scale(image)
83
+        results = self._predictor([scaled])
70
         if results:
84
         if results:
71
-            page = results[0]
72
-            for blk in getattr(page, "blocks", []):
73
-                html = getattr(blk, "html", "") or ""
74
-                label = getattr(blk, "label", "text")
75
-                bbox = getattr(blk, "bbox", [0, 0, 0, 0])
76
-                confidence = float(getattr(blk, "confidence", 0.9))
77
-                blocks.append(
78
-                    ParsedBlock(
79
-                        label=label,
80
-                        content=html,
81
-                        bbox=(int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3])),
82
-                        confidence=confidence,
83
-                    )
84
-                )
85
+            return self._to_result(results[0], scaled)
86
+        return OcrPageResult()
85
 
87
 
88
+    @staticmethod
89
+    def _to_result(page, image: PILImage.Image) -> OcrPageResult:
90
+        blocks: list[ParsedBlock] = []
91
+        for blk in getattr(page, "blocks", []):
92
+            html = getattr(blk, "html", "") or ""
93
+            label = getattr(blk, "label", "text")
94
+            bbox = getattr(blk, "bbox", [0, 0, 0, 0])
95
+            confidence = float(getattr(blk, "confidence", 0.9))
96
+            blocks.append(
97
+                ParsedBlock(
98
+                    label=label,
99
+                    content=html,
100
+                    bbox=(int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3])),
101
+                    confidence=confidence,
102
+                )
103
+            )
86
         return OcrPageResult(
104
         return OcrPageResult(
87
             blocks=blocks,
105
             blocks=blocks,
88
-            raw_json=raw_json,
106
+            raw_json={},
89
             width=image.width,
107
             width=image.width,
90
             height=image.height,
108
             height=image.height,
91
         )
109
         )
92
 
110
 
93
 
111
 
112
+def _pre_scale(image: PILImage.Image) -> PILImage.Image:
113
+    w, h = image.size
114
+    if w > _IMAGE_MAX_WIDTH:
115
+        ratio = _IMAGE_MAX_WIDTH / w
116
+        return image.resize((_IMAGE_MAX_WIDTH, int(h * ratio)), 1)  # PIL.Image.LANCZOS
117
+    return image
118
+
119
+
94
 _surya_engine: SuryaEngine | None = None
120
 _surya_engine: SuryaEngine | None = None
95
 
121
 
96
 
122
 
99
     if _surya_engine is None:
125
     if _surya_engine is None:
100
         _surya_engine = SuryaEngine()
126
         _surya_engine = SuryaEngine()
101
     return _surya_engine.process(image_path)
127
     return _surya_engine.process(image_path)
128
+
129
+
130
+def surya_ocr_batch(image_paths: Sequence[str]) -> list[OcrPageResult]:
131
+    global _surya_engine
132
+    if _surya_engine is None:
133
+        _surya_engine = SuryaEngine()
134
+    return _surya_engine.process_batch(image_paths)

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

9
 
9
 
10
 
10
 
11
 class Corrector(Protocol):
11
 class Corrector(Protocol):
12
-    def __call__(self, block: ParsedBlock) -> ParsedBlock: ...
12
+    def __call__(self, block: ParsedBlock) -> None: ...
13
 
13
 
14
 
14
 
15
 def apply_corrector(
15
 def apply_corrector(

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

7
 _SPACED_VARS = re.compile(r"\b([A-Z])\s+([A-Z])\b")
7
 _SPACED_VARS = re.compile(r"\b([A-Z])\s+([A-Z])\b")
8
 
8
 
9
 
9
 
10
-def fix_ocr_errors(blocks: list[ParsedBlock]) -> list[ParsedBlock]:
10
+def fix_ocr_errors(blocks: list[ParsedBlock]) -> None:
11
     for block in blocks:
11
     for block in blocks:
12
         if block.label == "formula":
12
         if block.label == "formula":
13
             block.content = _fix_formula_errors(block.content)
13
             block.content = _fix_formula_errors(block.content)
14
-    return blocks
15
 
14
 
16
 
15
 
17
 def _fix_formula_errors(content: str) -> str:
16
 def _fix_formula_errors(content: str) -> str:

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

18
             verbose=False,
18
             verbose=False,
19
         )
19
         )
20
 
20
 
21
-    def __call__(self, block: ParsedBlock) -> ParsedBlock:
22
-        if block.label == "formula":
23
-            return block
21
+    def __call__(self, block: ParsedBlock) -> None:
24
         if len(block.content) < 30:
22
         if len(block.content) < 30:
25
-            return block
23
+            return
26
 
24
 
27
         text = block.content[:400]
25
         text = block.content[:400]
28
         messages = [
26
         messages = [
34
 
32
 
35
         if corrected and len(corrected) > 10 and corrected != text:
33
         if corrected and len(corrected) > 10 and corrected != text:
36
             block.content = corrected
34
             block.content = corrected
37
-        return block
38
 
35
 
39
 
36
 
40
 _llm_corrector: LlmCorrector | None = None
37
 _llm_corrector: LlmCorrector | None = None

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

4
 import logging
4
 import logging
5
 import os
5
 import os
6
 import time as _time
6
 import time as _time
7
-from pathlib import Path
8
 from typing import Any
7
 from typing import Any
9
 
8
 
10
 import cv2
9
 import cv2
11
 import requests
10
 import requests
12
 
11
 
12
+from src.env import load_env
13
 from src.ocr.paddle_engine import ParsedBlock
13
 from src.ocr.paddle_engine import ParsedBlock
14
 
14
 
15
 logger = logging.getLogger(__name__)
15
 logger = logging.getLogger(__name__)
18
 API_URL_ENV: str = "OPENCODE_API_URL"
18
 API_URL_ENV: str = "OPENCODE_API_URL"
19
 DEFAULT_MODEL: str = "qwen3.8-max"
19
 DEFAULT_MODEL: str = "qwen3.8-max"
20
 
20
 
21
-_ENV_LOADED: bool = False
22
-
23
-
24
-def _load_env() -> None:
25
-    global _ENV_LOADED
26
-    if _ENV_LOADED:
27
-        return
28
-    env_path = Path(__file__).resolve().parent.parent.parent / ".env"
29
-    if env_path.exists():
30
-        with open(env_path) as f:
31
-            for line in f:
32
-                line = line.strip()
33
-                if line and not line.startswith("#") and "=" in line:
34
-                    key, _, value = line.partition("=")
35
-                    if key.strip() not in os.environ:
36
-                        os.environ[key.strip()] = value.strip().strip("\"'")
37
-    _ENV_LOADED = True
38
-
39
 
21
 
40
 def _get_api_key() -> str:
22
 def _get_api_key() -> str:
41
-    _load_env()
23
+    load_env()
42
     return os.environ.get(API_KEY_ENV, "")
24
     return os.environ.get(API_KEY_ENV, "")
43
 
25
 
44
 
26
 

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

176
 
176
 
177
 
177
 
178
 def find_content_bounds(image: np.ndarray) -> tuple[int, int, int, int]:
178
 def find_content_bounds(image: np.ndarray) -> tuple[int, int, int, int]:
179
-    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
180
-    blurred = cv2.GaussianBlur(gray, (3, 3), 0)
181
-    otsu_th, binary = cv2.threshold(
182
-        blurred, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU,
183
-    )
184
-    if otsu_th < 50:
185
-        binary = cv2.adaptiveThreshold(
186
-            blurred, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
187
-            cv2.THRESH_BINARY_INV, 31, 10,
188
-        )
189
-    kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (3, 3))
190
-    binary = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel)
179
+    binary = _binarize(image)
191
 
180
 
192
     h, w = binary.shape
181
     h, w = binary.shape
193
     min_density = int(w * DENSITY_RATIO)
182
     min_density = int(w * DENSITY_RATIO)

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