refs #62625 добавил прототип классификатора типов документов
This commit is contained in:
29
.gitignore
vendored
Normal file
29
.gitignore
vendored
Normal file
@@ -0,0 +1,29 @@
|
|||||||
|
# Виртуальные окружения
|
||||||
|
/.venv/
|
||||||
|
/.venv-document-classifier/
|
||||||
|
|
||||||
|
# Реальные документы, ссылки и датасеты
|
||||||
|
/dataset/
|
||||||
|
/test_dataset/
|
||||||
|
/downloads/
|
||||||
|
/links/
|
||||||
|
/tests/
|
||||||
|
|
||||||
|
# Обученные модели
|
||||||
|
/models/
|
||||||
|
*.pt
|
||||||
|
*.pth
|
||||||
|
*.onnx
|
||||||
|
|
||||||
|
# Служебные файлы Python и инструментов
|
||||||
|
__pycache__/
|
||||||
|
*.py[cod]
|
||||||
|
/.pytest_cache/
|
||||||
|
/.mypy_cache/
|
||||||
|
/.ruff_cache/
|
||||||
|
|
||||||
|
# Настройки IDE и ОС
|
||||||
|
/.idea/
|
||||||
|
/.vscode/
|
||||||
|
.DS_Store
|
||||||
|
Thumbs.db
|
||||||
300
README.md
Normal file
300
README.md
Normal file
@@ -0,0 +1,300 @@
|
|||||||
|
# Классификатор типов документов — задача 62625
|
||||||
|
|
||||||
|
Этот репозиторий — автономный тестовый прототип.
|
||||||
|
Он не импортирует код Vetro и не меняет текущий механизм загрузки. Сначала
|
||||||
|
модель необходимо обучить и измерить качество на документах, которых она не
|
||||||
|
видела.
|
||||||
|
|
||||||
|
## Текущее состояние модели
|
||||||
|
|
||||||
|
Текущая тестовая модель `request_document_types.pt` обучена распознавать только
|
||||||
|
четыре типа документов:
|
||||||
|
|
||||||
|
| Значение `RequestFileType` | Тип документа |
|
||||||
|
|---|---|
|
||||||
|
| `referral_for_repairs` | Направление на ремонт |
|
||||||
|
| `inspection_act` | Акт осмотра |
|
||||||
|
| `cc_approval` | Согласование страхового случая |
|
||||||
|
| `inspection_photo` | Фото осмотра |
|
||||||
|
|
||||||
|
Документ другого типа модель пока не сможет корректно определить: она выберет
|
||||||
|
наиболее похожий из этих четырёх классов либо вернёт `unknown`, если
|
||||||
|
уверенность ниже установленного порога.
|
||||||
|
|
||||||
|
### Полученные результаты распознавания
|
||||||
|
|
||||||
|
Порог принятия результата — 0.70. Все четыре проверенных документа были
|
||||||
|
приняты моделью:
|
||||||
|
|
||||||
|
| Файл | Распознанный тип | Уверенность | Следующие варианты |
|
||||||
|
|---|---|---:|---|
|
||||||
|
| `1.pdf` | `referral_for_repairs` | 99.6794% | `cc_approval` — 0.2871%, `inspection_act` — 0.0221% |
|
||||||
|
| `2.pdf` | `inspection_act` | 97.1619% | `referral_for_repairs` — 2.2419%, `cc_approval` — 0.3883% |
|
||||||
|
| `3.png` | `cc_approval` | 92.1763% | `referral_for_repairs` — 7.6405%, `inspection_act` — 0.1245% |
|
||||||
|
| `4.jpg` | `inspection_photo` | 99.7211% | `inspection_act` — 0.2390%, `referral_for_repairs` — 0.0324% |
|
||||||
|
|
||||||
|
Каждый файл содержал одну страницу. Эти результаты подтверждают, что запуск
|
||||||
|
инференса работает и модель уверенно классифицировала четыре конкретных
|
||||||
|
примера. Они не являются оценкой общей точности модели: для такой оценки
|
||||||
|
нужно выполнить команду `evaluate` на независимом `test_dataset`, который не
|
||||||
|
использовался при обучении.
|
||||||
|
|
||||||
|
## 1. Подготовка окружения
|
||||||
|
|
||||||
|
Для Windows рекомендуется Python 3.12 (64-bit) и отдельное виртуальное
|
||||||
|
окружение. PyTorch для Windows официально поддерживает Python до версии 3.12.
|
||||||
|
Если Python ещё не установлен, используйте Python 3.12.10 — это последняя
|
||||||
|
версия ветки 3.12 с готовым Windows installer.
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python -m venv .venv-document-classifier
|
||||||
|
.\.venv-document-classifier\Scripts\Activate.ps1
|
||||||
|
python -m pip install --upgrade pip
|
||||||
|
python -m pip install -r requirements.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
### Если существующее окружение перестало запускаться
|
||||||
|
|
||||||
|
Файл `.venv-document-classifier\pyvenv.cfg` содержит путь к Python, на основе
|
||||||
|
которого создавалось окружение. Если этот Python был удалён или перемещён,
|
||||||
|
появляется ошибка `Unable to create process`. Такое окружение нужно создать
|
||||||
|
заново после установки Python 3.12:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
deactivate
|
||||||
|
Rename-Item .venv-document-classifier .venv-document-classifier-old
|
||||||
|
py -3.12 -m venv .venv-document-classifier
|
||||||
|
.\.venv-document-classifier\Scripts\Activate.ps1
|
||||||
|
python -m pip install --upgrade pip
|
||||||
|
python -m pip install -r requirements.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
Старый каталог оставлен как резервная копия. После успешной проверки новой
|
||||||
|
среды его можно удалить вручную.
|
||||||
|
|
||||||
|
В `requirements.txt` зафиксированы совместимые версии PyTorch и torchvision
|
||||||
|
с CUDA 12.8 и официальный индекс `cu128`. Поэтому все зависимости устанавливаются
|
||||||
|
одной командой. Обычный `pip install torch torchvision` без CUDA-индекса может
|
||||||
|
установить CPU-сборку, которая не умеет работать с RTX 3060 Ti.
|
||||||
|
|
||||||
|
Если CPU-сборка уже установлена, исправьте текущее активное окружение:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python -m pip uninstall -y torch torchvision
|
||||||
|
python -m pip install --no-cache-dir -r requirements.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
Проверьте, что PyTorch видит видеокарту:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python -c "import torch; ok=torch.cuda.is_available(); print('torch:',torch.__version__,'cuda build:',torch.version.cuda,'available:',ok,'GPU:',torch.cuda.get_device_name(0) if ok else 'нет')"
|
||||||
|
```
|
||||||
|
|
||||||
|
Ожидаемый результат на текущем ПК: `True` и `NVIDIA GeForce RTX 3060 Ti`.
|
||||||
|
Если выводится `False`, установлена CPU-сборка PyTorch — до обучения нужно
|
||||||
|
поставить CUDA-сборку командой выше.
|
||||||
|
|
||||||
|
Для первого запуска обучения нужен интернет: torchvision скачает
|
||||||
|
предобученные ImageNet-веса MobileNetV3. Если интернета нет, можно добавить
|
||||||
|
`--no-pretrained`, но качество обычно будет ниже.
|
||||||
|
|
||||||
|
## 2. Подготовка датасета
|
||||||
|
|
||||||
|
Создайте один каталог на каждый тип. Имя каталога должно быть **точным
|
||||||
|
строковым значением** `RequestFileType` из Vetro, а не русским заголовком:
|
||||||
|
|
||||||
|
```text
|
||||||
|
dataset/
|
||||||
|
├── referral_for_repairs/
|
||||||
|
│ ├── direction_0001.pdf
|
||||||
|
│ └── direction_0002.jpg
|
||||||
|
├── inspection_act/
|
||||||
|
│ ├── act_0001.pdf
|
||||||
|
│ └── act_0002.png
|
||||||
|
└── check/
|
||||||
|
├── invoice_0001.pdf
|
||||||
|
└── invoice_0002.pdf
|
||||||
|
```
|
||||||
|
|
||||||
|
Поддерживаются PDF, PNG, JPG/JPEG, BMP, TIFF и WEBP. DOC/DOCX/XLS следует
|
||||||
|
заранее преобразовать в PDF. Один многостраничный PDF считается одним
|
||||||
|
документом: все его страницы попадут только в train или только в validation.
|
||||||
|
|
||||||
|
Практические требования к данным:
|
||||||
|
|
||||||
|
- минимум — 2 документа на тип, разумный старт — 50–100;
|
||||||
|
- лучше 200+ документов на каждый часто встречающийся тип;
|
||||||
|
- не складывайте копии одного документа под разными именами;
|
||||||
|
- включите сканы и фото разного качества, документы разных страховых компаний;
|
||||||
|
- номера дел, ФИО и прочие персональные данные должны храниться только в
|
||||||
|
защищённой инфраструктуре; этот скрипт никуда их не отправляет;
|
||||||
|
- число документов по типам желательно выровнять. В коде есть компенсация
|
||||||
|
дисбаланса классов, но она не заменяет реальные примеры.
|
||||||
|
|
||||||
|
Отложите отдельный каталог `test_dataset` (обычно 10–20% исходных документов)
|
||||||
|
**до обучения**. Он должен иметь ту же структуру и не должен пересекаться с
|
||||||
|
`dataset`. Встроенный validation нужен для выбора эпохи, а `test_dataset` —
|
||||||
|
для честной финальной оценки.
|
||||||
|
|
||||||
|
## 3. Обучение
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python document_type_classifier.py train `
|
||||||
|
--data .\dataset `
|
||||||
|
--model .\models\request_document_types.pt `
|
||||||
|
--device cuda `
|
||||||
|
--epochs 25 `
|
||||||
|
--batch-size 64 `
|
||||||
|
--workers 0 `
|
||||||
|
--log-every 10
|
||||||
|
```
|
||||||
|
|
||||||
|
Эти значения подобраны для текущего компьютера:
|
||||||
|
|
||||||
|
- AMD Ryzen 5 7500F: 6 ядер / 12 потоков;
|
||||||
|
- 32 ГБ оперативной памяти;
|
||||||
|
- NVIDIA GeForce RTX 3060 Ti, 8 ГБ видеопамяти.
|
||||||
|
|
||||||
|
По умолчанию уже используются `batch-size=64`, `workers=0`, размер изображения
|
||||||
|
224×224 и mixed precision. В памяти хранится только индекс путей и страниц;
|
||||||
|
каждое изображение декодируется непосредственно перед своим batch. Это не
|
||||||
|
позволяет фотографиям высокого разрешения занять всю оперативную память перед
|
||||||
|
первой эпохой. Явные параметры в примере оставлены для воспроизводимости. Если
|
||||||
|
появляется `CUDA out of memory`, сначала установите
|
||||||
|
`--batch-size 32`, затем 16. Если CUDA работает нестабильно, добавьте
|
||||||
|
`--no-amp`; для принудительного CPU используйте `--device cpu --batch-size 16
|
||||||
|
--workers 0`.
|
||||||
|
|
||||||
|
Во время запуска логируются сканирование и состав датасета, чтение каждых 25
|
||||||
|
документов, загрузка модели, каждый 10-й batch, validation, время эпохи и
|
||||||
|
примерное оставшееся время. Частоту можно изменить параметрами
|
||||||
|
`--dataset-log-every` и `--log-every`.
|
||||||
|
|
||||||
|
Скрипт печатает `val_accuracy` после каждой эпохи, сохраняет только лучшую
|
||||||
|
модель и останавливается после пяти эпох без улучшения.
|
||||||
|
|
||||||
|
Реальные JPEG из Vetro могут содержать неполный последний блок или
|
||||||
|
нестандартную структуру, хотя нормально открываются браузером. Для таких
|
||||||
|
файлов включён tolerant-режим Pillow: доступная часть изображения участвует в
|
||||||
|
обучении. Сообщение `Пропуск поврежденного файла` теперь остаётся только для
|
||||||
|
файлов, которые действительно невозможно декодировать.
|
||||||
|
|
||||||
|
## 4. Честная проверка
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python document_type_classifier.py evaluate `
|
||||||
|
--model .\models\request_document_types.pt `
|
||||||
|
--data .\test_dataset
|
||||||
|
```
|
||||||
|
|
||||||
|
В отчёте:
|
||||||
|
|
||||||
|
- `accuracy_with_unknown_as_error` — доля всех верных ответов;
|
||||||
|
- `accepted_share` — доля документов, для которых уверенность не ниже порога;
|
||||||
|
- `by_expected_type` — ошибки по каждому реальному типу.
|
||||||
|
|
||||||
|
Для автозагрузки недостаточно одной общей accuracy. Проверьте каждый тип и
|
||||||
|
особенно пары, которые модель путает. Рекомендуемый критерий пилота: не менее
|
||||||
|
95% точности среди принятых ответов; остальные файлы должны оставаться на
|
||||||
|
ручном выборе типа.
|
||||||
|
|
||||||
|
Порог можно сделать строже:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python document_type_classifier.py evaluate `
|
||||||
|
--model .\models\request_document_types.pt `
|
||||||
|
--data .\test_dataset `
|
||||||
|
--threshold 0.85
|
||||||
|
```
|
||||||
|
|
||||||
|
## 5. Распознавание
|
||||||
|
|
||||||
|
Один документ:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python document_type_classifier.py predict `
|
||||||
|
--model .\models\request_document_types.pt `
|
||||||
|
--file .\document.pdf
|
||||||
|
```
|
||||||
|
|
||||||
|
Каталог документов:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python document_type_classifier.py predict `
|
||||||
|
--model .\models\request_document_types.pt `
|
||||||
|
--file .\incoming
|
||||||
|
```
|
||||||
|
|
||||||
|
На каждый файл выводится одна JSON-строка:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{"file":"document.pdf","document_type":"inspection_act","confidence":0.9631,"accepted":true,"pages":2,"alternatives":[...]}
|
||||||
|
```
|
||||||
|
|
||||||
|
Если уверенность ниже 0.70, `document_type` будет `unknown` и
|
||||||
|
`accepted=false`. При интеграции такой документ нужно показывать пользователю
|
||||||
|
для ручного выбора. `alternatives` полезны для интерфейса и анализа ошибок.
|
||||||
|
|
||||||
|
Все команды следует выполнять из корня репозитория.
|
||||||
|
|
||||||
|
## 6. Возможная интеграция после завершения тестового пилота
|
||||||
|
|
||||||
|
1. Загрузить модель один раз при старте отдельного worker/API-сервиса, а не
|
||||||
|
заново для каждого HTTP-запроса.
|
||||||
|
2. После получения временного файла вызвать классификацию до
|
||||||
|
`document.upload_file(...)`.
|
||||||
|
3. Проверить, что предсказанная строка входит в `RequestFileType`.
|
||||||
|
4. При `accepted=true` передать её как `filetype`; при `false` сохранить
|
||||||
|
текущий ручной выбор.
|
||||||
|
5. Логировать ожидаемый/исправленный пользователем тип и confidence. Эти
|
||||||
|
исправления формируют следующую версию датасета.
|
||||||
|
6. Версионировать файл модели и иметь возможность мгновенно выключить
|
||||||
|
автоклассификацию конфигурационным флагом.
|
||||||
|
|
||||||
|
Не рекомендуется запускать тяжёлый инференс синхронно внутри нескольких
|
||||||
|
web-процессов без замера времени и памяти. Для промышленной интеграции лучше
|
||||||
|
обернуть классификатор внутренним API или фоновой задачей и ограничить размер
|
||||||
|
и число страниц входного файла.
|
||||||
|
|
||||||
|
## 7. Скачивание документов через авторизованный браузер
|
||||||
|
|
||||||
|
Создайте, например, `links.txt`, указав одну HTTP/HTTPS-ссылку на строку:
|
||||||
|
|
||||||
|
```text
|
||||||
|
https://example.org/document-1.pdf
|
||||||
|
https://example.org/document-2.jpg
|
||||||
|
```
|
||||||
|
|
||||||
|
Скрипт не скачивает файлы HTTP-клиентом Python. Он последовательно передаёт
|
||||||
|
ссылки системному браузеру, поэтому используется уже существующая браузерная
|
||||||
|
авторизация, а имя и расширение файла определяет сам сервер.
|
||||||
|
|
||||||
|
Перед запуском откройте настройки загрузок Chrome/Edge:
|
||||||
|
|
||||||
|
1. Установите каталог загрузки:
|
||||||
|
`.\downloads`.
|
||||||
|
2. Отключите запрос места сохранения для каждого файла.
|
||||||
|
3. Разрешите сайту автоматическое скачивание нескольких файлов, если браузер
|
||||||
|
покажет соответствующий запрос.
|
||||||
|
|
||||||
|
Затем запустите:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python download_files.py links.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
Пустые строки и строки, начинающиеся с `#`, игнорируются. Между ссылками по
|
||||||
|
умолчанию выдерживается пауза 1.5 секунды. Для пробного запуска первых десяти:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python download_files.py links.txt --limit 10
|
||||||
|
```
|
||||||
|
|
||||||
|
Если выполнение было прервано после 350-й ссылки:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python download_files.py links.txt --start 351
|
||||||
|
```
|
||||||
|
|
||||||
|
Пауза настраивается параметром `--delay`, например `--delay 3`. Скрипт не
|
||||||
|
переименовывает файлы: браузер сохраняет оригинальное имя из ответа сервера.
|
||||||
697
document_type_classifier.py
Normal file
697
document_type_classifier.py
Normal file
@@ -0,0 +1,697 @@
|
|||||||
|
"""
|
||||||
|
Обучение и запуск нейросетевого классификатора типов документов.
|
||||||
|
|
||||||
|
Поддерживаются PDF, PNG, JPG/JPEG, BMP, TIFF и WEBP. Один каталог внутри
|
||||||
|
датасета соответствует одному значению RequestFileType, например:
|
||||||
|
|
||||||
|
dataset/
|
||||||
|
referral_for_repairs/
|
||||||
|
direction_001.pdf
|
||||||
|
inspection_act/
|
||||||
|
act_001.jpg
|
||||||
|
|
||||||
|
Примеры:
|
||||||
|
python document_type_classifier.py train --data dataset --model models/document_types.pt
|
||||||
|
python document_type_classifier.py predict --model models/document_types.pt --file example.pdf
|
||||||
|
python document_type_classifier.py evaluate --model models/document_types.pt --data test_dataset
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import random
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from collections import Counter
|
||||||
|
from dataclasses import asdict, dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Iterable
|
||||||
|
|
||||||
|
SUPPORTED_EXTENSIONS = {".pdf", ".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff", ".webp"}
|
||||||
|
DEFAULT_IMAGE_SIZE = 224
|
||||||
|
DEFAULT_CONFIDENCE_THRESHOLD = 0.70
|
||||||
|
DEFAULT_BATCH_SIZE = 64
|
||||||
|
DEFAULT_WORKERS = 0
|
||||||
|
DEFAULT_LOG_EVERY = 10
|
||||||
|
DEFAULT_DATASET_LOG_EVERY = 25
|
||||||
|
|
||||||
|
|
||||||
|
def _dependencies() -> tuple[Any, Any, Any, Any, Any]:
|
||||||
|
"""Import heavy optional dependencies only when a command needs them."""
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
from PIL import Image, ImageFile
|
||||||
|
from torch import nn
|
||||||
|
from torch.utils.data import DataLoader, Dataset
|
||||||
|
from torchvision import models, transforms
|
||||||
|
except ImportError as exc:
|
||||||
|
raise SystemExit(
|
||||||
|
"Не установлены зависимости. Выполните:\n"
|
||||||
|
"pip install torch torchvision pillow pypdfium2"
|
||||||
|
) from exc
|
||||||
|
# Реальные фотографии из Vetro иногда имеют корректную JPEG-сигнатуру,
|
||||||
|
# но содержат нестандартный или неполный последний блок. Браузеры и
|
||||||
|
# просмотрщики открывают их, поэтому разрешаем Pillow дочитать доступные
|
||||||
|
# пиксели вместо исключения "image file is truncated / No data for frame".
|
||||||
|
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||||
|
return torch, Image, nn, (DataLoader, Dataset), (models, transforms)
|
||||||
|
|
||||||
|
|
||||||
|
def _render_document(path: Path, dpi_scale: float = 1.5) -> list[Any]:
|
||||||
|
"""Return RGB PIL images, one for every page/frame of a document."""
|
||||||
|
_, Image, _, _, _ = _dependencies()
|
||||||
|
if path.suffix.lower() == ".pdf":
|
||||||
|
try:
|
||||||
|
import pypdfium2 as pdfium
|
||||||
|
except ImportError as exc:
|
||||||
|
raise SystemExit("Для PDF установите pypdfium2: pip install pypdfium2") from exc
|
||||||
|
|
||||||
|
pdf = pdfium.PdfDocument(str(path))
|
||||||
|
try:
|
||||||
|
return [page.render(scale=dpi_scale).to_pil().convert("RGB") for page in pdf]
|
||||||
|
finally:
|
||||||
|
pdf.close()
|
||||||
|
|
||||||
|
image = Image.open(path)
|
||||||
|
try:
|
||||||
|
# У некоторых реальных JPEG из Vetro повреждена только служебная
|
||||||
|
# таблица кадров: пиксели читаются, но даже seek(0) завершается ошибкой
|
||||||
|
# "No data found for frame". Кадровая навигация нужна только TIFF.
|
||||||
|
if path.suffix.lower() not in {".tif", ".tiff"}:
|
||||||
|
return [image.convert("RGB").copy()]
|
||||||
|
|
||||||
|
pages = []
|
||||||
|
frame_count = getattr(image, "n_frames", 1)
|
||||||
|
for frame_index in range(frame_count):
|
||||||
|
image.seek(frame_index)
|
||||||
|
pages.append(image.convert("RGB").copy())
|
||||||
|
return pages
|
||||||
|
finally:
|
||||||
|
image.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _document_page_count(path: Path) -> int:
|
||||||
|
"""Вернуть число страниц/кадров, не сохраняя изображение в памяти."""
|
||||||
|
_, Image, _, _, _ = _dependencies()
|
||||||
|
if path.suffix.lower() == ".pdf":
|
||||||
|
try:
|
||||||
|
import pypdfium2 as pdfium
|
||||||
|
except ImportError as exc:
|
||||||
|
raise SystemExit("Для PDF установите pypdfium2: pip install pypdfium2") from exc
|
||||||
|
|
||||||
|
pdf = pdfium.PdfDocument(str(path))
|
||||||
|
try:
|
||||||
|
return len(pdf)
|
||||||
|
finally:
|
||||||
|
pdf.close()
|
||||||
|
|
||||||
|
image = Image.open(path)
|
||||||
|
try:
|
||||||
|
if path.suffix.lower() in {".tif", ".tiff"}:
|
||||||
|
return getattr(image, "n_frames", 1)
|
||||||
|
return 1
|
||||||
|
finally:
|
||||||
|
image.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _render_page(path: Path, page_index: int, dpi_scale: float = 1.5) -> Any:
|
||||||
|
"""Декодировать одну страницу документа непосредственно перед batch."""
|
||||||
|
_, Image, _, _, _ = _dependencies()
|
||||||
|
if path.suffix.lower() == ".pdf":
|
||||||
|
try:
|
||||||
|
import pypdfium2 as pdfium
|
||||||
|
except ImportError as exc:
|
||||||
|
raise SystemExit("Для PDF установите pypdfium2: pip install pypdfium2") from exc
|
||||||
|
|
||||||
|
pdf = pdfium.PdfDocument(str(path))
|
||||||
|
page = pdf[page_index]
|
||||||
|
try:
|
||||||
|
return page.render(scale=dpi_scale).to_pil().convert("RGB")
|
||||||
|
finally:
|
||||||
|
page.close()
|
||||||
|
pdf.close()
|
||||||
|
|
||||||
|
image = Image.open(path)
|
||||||
|
try:
|
||||||
|
if path.suffix.lower() in {".tif", ".tiff"}:
|
||||||
|
image.seek(page_index)
|
||||||
|
return image.convert("RGB").copy()
|
||||||
|
finally:
|
||||||
|
image.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _document_files(directory: Path) -> list[Path]:
|
||||||
|
return sorted(
|
||||||
|
path
|
||||||
|
for path in directory.rglob("*")
|
||||||
|
if path.is_file() and path.suffix.lower() in SUPPORTED_EXTENSIONS
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _scan_dataset(root: Path) -> dict[str, list[Path]]:
|
||||||
|
print(f"[Датасет] Сканирование каталога: {root.resolve()}", flush=True)
|
||||||
|
if not root.is_dir():
|
||||||
|
raise SystemExit(f"Каталог датасета не найден: {root}")
|
||||||
|
classes = {
|
||||||
|
child.name: _document_files(child)
|
||||||
|
for child in sorted(root.iterdir())
|
||||||
|
if child.is_dir()
|
||||||
|
}
|
||||||
|
classes = {label: paths for label, paths in classes.items() if paths}
|
||||||
|
if len(classes) < 2:
|
||||||
|
raise SystemExit("В датасете должны быть хотя бы два непустых каталога типов документов.")
|
||||||
|
print(
|
||||||
|
f"[Датасет] Найдено типов: {len(classes)}, документов: "
|
||||||
|
f"{sum(len(paths) for paths in classes.values())}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
for label, paths in classes.items():
|
||||||
|
print(f" {label}: {len(paths)}", flush=True)
|
||||||
|
return classes
|
||||||
|
|
||||||
|
|
||||||
|
def _split_documents(
|
||||||
|
classes: dict[str, list[Path]], validation_share: float, seed: int
|
||||||
|
) -> tuple[list[tuple[Path, int]], list[tuple[Path, int]], list[str]]:
|
||||||
|
labels = sorted(classes)
|
||||||
|
train: list[tuple[Path, int]] = []
|
||||||
|
validation: list[tuple[Path, int]] = []
|
||||||
|
rng = random.Random(seed)
|
||||||
|
|
||||||
|
for class_index, label in enumerate(labels):
|
||||||
|
paths = classes[label][:]
|
||||||
|
if len(paths) < 2:
|
||||||
|
raise SystemExit(
|
||||||
|
f"Для типа '{label}' нужен минимум 2 документа, рекомендуется не менее 50."
|
||||||
|
)
|
||||||
|
rng.shuffle(paths)
|
||||||
|
validation_count = max(1, round(len(paths) * validation_share))
|
||||||
|
validation_count = min(validation_count, len(paths) - 1)
|
||||||
|
validation.extend((path, class_index) for path in paths[:validation_count])
|
||||||
|
train.extend((path, class_index) for path in paths[validation_count:])
|
||||||
|
rng.shuffle(train)
|
||||||
|
rng.shuffle(validation)
|
||||||
|
return train, validation, labels
|
||||||
|
|
||||||
|
|
||||||
|
def _transforms(image_size: int, training: bool) -> Any:
|
||||||
|
_, _, _, _, (_, transforms) = _dependencies()
|
||||||
|
operations: list[Any] = [
|
||||||
|
transforms.Resize((image_size, image_size)),
|
||||||
|
]
|
||||||
|
if training:
|
||||||
|
# Small distortions imitate phone photos without mirroring document text.
|
||||||
|
operations.extend(
|
||||||
|
[
|
||||||
|
transforms.RandomRotation(3),
|
||||||
|
transforms.ColorJitter(brightness=0.15, contrast=0.15),
|
||||||
|
transforms.RandomPerspective(distortion_scale=0.08, p=0.25),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
operations.extend(
|
||||||
|
[
|
||||||
|
transforms.ToTensor(),
|
||||||
|
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
return transforms.Compose(operations)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_model(class_count: int, pretrained: bool) -> Any:
|
||||||
|
_, _, nn, _, (models, _) = _dependencies()
|
||||||
|
weights = models.MobileNet_V3_Small_Weights.DEFAULT if pretrained else None
|
||||||
|
model = models.mobilenet_v3_small(weights=weights)
|
||||||
|
input_features = model.classifier[-1].in_features
|
||||||
|
model.classifier[-1] = nn.Linear(input_features, class_count)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def _device(torch: Any, requested: str) -> Any:
|
||||||
|
if requested == "auto":
|
||||||
|
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
if requested == "cuda" and not torch.cuda.is_available():
|
||||||
|
raise SystemExit("CUDA запрошена, но недоступна.")
|
||||||
|
return torch.device(requested)
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_cuda(torch: Any, device: Any) -> None:
|
||||||
|
"""Fast settings suitable for the local RTX 3060 Ti."""
|
||||||
|
if device.type != "cuda":
|
||||||
|
return
|
||||||
|
torch.backends.cudnn.benchmark = True
|
||||||
|
# Ampere supports TF32; it speeds up float32 operations with negligible
|
||||||
|
# influence on a document classification baseline.
|
||||||
|
torch.backends.cuda.matmul.allow_tf32 = True
|
||||||
|
torch.backends.cudnn.allow_tf32 = True
|
||||||
|
|
||||||
|
|
||||||
|
def _document_dataset(base_class: Any) -> type:
|
||||||
|
class DocumentPageDataset(base_class):
|
||||||
|
"""Набор страниц с ленивым декодированием изображений."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
documents: list[tuple[Path, int]],
|
||||||
|
transform: Any,
|
||||||
|
phase: str,
|
||||||
|
log_every: int,
|
||||||
|
):
|
||||||
|
self.items: list[tuple[Path, int, int]] = []
|
||||||
|
self.transform = transform
|
||||||
|
started_at = time.monotonic()
|
||||||
|
skipped = 0
|
||||||
|
print(f"[{phase}] Индексация документов: {len(documents)}", flush=True)
|
||||||
|
for number, (path, label) in enumerate(documents, start=1):
|
||||||
|
try:
|
||||||
|
page_count = _document_page_count(path)
|
||||||
|
if page_count < 1:
|
||||||
|
raise ValueError("документ не содержит страниц")
|
||||||
|
except Exception as exc:
|
||||||
|
skipped += 1
|
||||||
|
print(
|
||||||
|
f"[{phase}] Пропуск поврежденного файла {path}: {exc}",
|
||||||
|
file=sys.stderr,
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
self.items.extend((path, label, page_index) for page_index in range(page_count))
|
||||||
|
if number % log_every == 0 or number == len(documents):
|
||||||
|
elapsed = time.monotonic() - started_at
|
||||||
|
speed = number / elapsed if elapsed else 0
|
||||||
|
print(
|
||||||
|
f"[{phase}] {number}/{len(documents)} документов, "
|
||||||
|
f"страниц: {len(self.items)}, пропущено: {skipped}, "
|
||||||
|
f"{speed:.1f} док/с",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[{phase}] Подготовка завершена за {time.monotonic() - started_at:.1f} с. "
|
||||||
|
f"Страниц: {len(self.items)}, пропущено документов: {skipped}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.items)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> tuple[Any, int, str]:
|
||||||
|
path, label, page_index = self.items[index]
|
||||||
|
try:
|
||||||
|
image = _render_page(path, page_index)
|
||||||
|
except Exception as exc:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Не удалось декодировать страницу {page_index + 1} файла {path}: {exc}"
|
||||||
|
) from exc
|
||||||
|
return self.transform(image), label, str(path)
|
||||||
|
|
||||||
|
return DocumentPageDataset
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Prediction:
|
||||||
|
file: str
|
||||||
|
document_type: str
|
||||||
|
confidence: float
|
||||||
|
accepted: bool
|
||||||
|
pages: int
|
||||||
|
alternatives: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
def _load_checkpoint(model_path: Path, device: Any) -> tuple[Any, dict[str, Any]]:
|
||||||
|
torch, _, _, _, _ = _dependencies()
|
||||||
|
try:
|
||||||
|
checkpoint = torch.load(model_path, map_location=device, weights_only=True)
|
||||||
|
except TypeError: # Compatibility with older PyTorch.
|
||||||
|
checkpoint = torch.load(model_path, map_location=device)
|
||||||
|
required = {"state_dict", "labels", "image_size"}
|
||||||
|
if not required.issubset(checkpoint):
|
||||||
|
raise SystemExit(f"Некорректный файл модели: {model_path}")
|
||||||
|
model = _make_model(len(checkpoint["labels"]), pretrained=False)
|
||||||
|
model.load_state_dict(checkpoint["state_dict"])
|
||||||
|
model.to(device)
|
||||||
|
model.eval()
|
||||||
|
return model, checkpoint
|
||||||
|
|
||||||
|
|
||||||
|
def _predict(
|
||||||
|
model: Any,
|
||||||
|
labels: list[str],
|
||||||
|
image_size: int,
|
||||||
|
path: Path,
|
||||||
|
device: Any,
|
||||||
|
threshold: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Prediction:
|
||||||
|
torch, _, _, _, _ = _dependencies()
|
||||||
|
pages = _render_document(path)
|
||||||
|
if not pages:
|
||||||
|
raise ValueError("Документ не содержит страниц")
|
||||||
|
transform = _transforms(image_size, training=False)
|
||||||
|
probability_sum = torch.zeros(len(labels), device=device)
|
||||||
|
inference_batch_size = DEFAULT_BATCH_SIZE if device.type == "cuda" else 16
|
||||||
|
with torch.inference_mode():
|
||||||
|
for start in range(0, len(pages), inference_batch_size):
|
||||||
|
batch = torch.stack(
|
||||||
|
[transform(page) for page in pages[start : start + inference_batch_size]]
|
||||||
|
).to(device, non_blocking=device.type == "cuda")
|
||||||
|
with torch.autocast(device_type=device.type, enabled=device.type == "cuda"):
|
||||||
|
probability_sum += torch.softmax(model(batch), dim=1).sum(dim=0)
|
||||||
|
# Mean of page probabilities gives one prediction for a multi-page document.
|
||||||
|
probabilities = probability_sum / len(pages)
|
||||||
|
count = min(max(1, top_k), len(labels))
|
||||||
|
scores, indices = torch.topk(probabilities, count)
|
||||||
|
alternatives = [
|
||||||
|
{"document_type": labels[index], "confidence": round(float(score), 6)}
|
||||||
|
for score, index in zip(scores.cpu().tolist(), indices.cpu().tolist())
|
||||||
|
]
|
||||||
|
confidence = float(scores[0])
|
||||||
|
return Prediction(
|
||||||
|
file=str(path),
|
||||||
|
document_type=labels[int(indices[0])] if confidence >= threshold else "unknown",
|
||||||
|
confidence=round(confidence, 6),
|
||||||
|
accepted=confidence >= threshold,
|
||||||
|
pages=len(pages),
|
||||||
|
alternatives=alternatives,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def train(args: argparse.Namespace) -> None:
|
||||||
|
torch, _, nn, (DataLoader, Dataset), _ = _dependencies()
|
||||||
|
started_at = time.monotonic()
|
||||||
|
print("[Запуск] Инициализация обучения", flush=True)
|
||||||
|
random.seed(args.seed)
|
||||||
|
torch.manual_seed(args.seed)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.manual_seed_all(args.seed)
|
||||||
|
|
||||||
|
classes = _scan_dataset(args.data)
|
||||||
|
train_documents, validation_documents, labels = _split_documents(
|
||||||
|
classes, args.validation_share, args.seed
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[Датасет] Разбиение по документам: train={len(train_documents)}, "
|
||||||
|
f"validation={len(validation_documents)}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
DatasetClass = _document_dataset(Dataset)
|
||||||
|
train_dataset = DatasetClass(
|
||||||
|
train_documents,
|
||||||
|
_transforms(args.image_size, training=True),
|
||||||
|
phase="Train dataset",
|
||||||
|
log_every=args.dataset_log_every,
|
||||||
|
)
|
||||||
|
validation_dataset = DatasetClass(
|
||||||
|
validation_documents,
|
||||||
|
_transforms(args.image_size, training=False),
|
||||||
|
phase="Validation dataset",
|
||||||
|
log_every=args.dataset_log_every,
|
||||||
|
)
|
||||||
|
if not train_dataset or not validation_dataset:
|
||||||
|
raise SystemExit("После чтения файлов обучающая или проверочная выборка пуста.")
|
||||||
|
|
||||||
|
device = _device(torch, args.device)
|
||||||
|
_configure_cuda(torch, device)
|
||||||
|
if device.type == "cuda":
|
||||||
|
properties = torch.cuda.get_device_properties(device)
|
||||||
|
print(
|
||||||
|
f"[Устройство] {properties.name}, VRAM: "
|
||||||
|
f"{properties.total_memory / 1024**3:.1f} ГБ, AMP: {not args.no_amp}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
print("[Устройство] CPU", flush=True)
|
||||||
|
print(
|
||||||
|
"[Модель] Загрузка MobileNetV3 и предобученных весов"
|
||||||
|
if not args.no_pretrained
|
||||||
|
else "[Модель] Создание MobileNetV3 без предобученных весов",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
model = _make_model(len(labels), pretrained=not args.no_pretrained).to(device)
|
||||||
|
if device.type == "cuda":
|
||||||
|
model = model.to(memory_format=torch.channels_last)
|
||||||
|
train_counts = Counter(label for _, label, _ in train_dataset.items)
|
||||||
|
class_weights = torch.tensor(
|
||||||
|
[len(train_dataset) / (len(labels) * train_counts[index]) for index in range(len(labels))],
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
criterion = nn.CrossEntropyLoss(weight=class_weights)
|
||||||
|
optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=1e-4)
|
||||||
|
use_amp = device.type == "cuda" and not args.no_amp
|
||||||
|
scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
|
||||||
|
train_loader = DataLoader(
|
||||||
|
train_dataset,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
shuffle=True,
|
||||||
|
num_workers=args.workers,
|
||||||
|
pin_memory=device.type == "cuda",
|
||||||
|
)
|
||||||
|
validation_loader = DataLoader(
|
||||||
|
validation_dataset,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
shuffle=False,
|
||||||
|
num_workers=args.workers,
|
||||||
|
pin_memory=device.type == "cuda",
|
||||||
|
)
|
||||||
|
|
||||||
|
best_accuracy = -1.0
|
||||||
|
epochs_without_improvement = 0
|
||||||
|
args.model.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
print(
|
||||||
|
f"Устройство: {device}; типов: {len(labels)}; "
|
||||||
|
f"страниц train/val: {len(train_dataset)}/{len(validation_dataset)}; "
|
||||||
|
f"batch size: {args.batch_size}; workers: {args.workers}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
for epoch in range(1, args.epochs + 1):
|
||||||
|
epoch_started_at = time.monotonic()
|
||||||
|
model.train()
|
||||||
|
running_loss = 0.0
|
||||||
|
print(
|
||||||
|
f"[Эпоха {epoch}/{args.epochs}] Обучение, batches: {len(train_loader)}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
for batch_number, (images, targets, _) in enumerate(train_loader, start=1):
|
||||||
|
images = images.to(device, non_blocking=device.type == "cuda")
|
||||||
|
targets = targets.to(device, non_blocking=device.type == "cuda")
|
||||||
|
if device.type == "cuda":
|
||||||
|
images = images.to(memory_format=torch.channels_last)
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
with torch.autocast(device_type=device.type, enabled=use_amp):
|
||||||
|
loss = criterion(model(images), targets)
|
||||||
|
scaler.scale(loss).backward()
|
||||||
|
scaler.step(optimizer)
|
||||||
|
scaler.update()
|
||||||
|
running_loss += float(loss) * images.size(0)
|
||||||
|
if batch_number % args.log_every == 0 or batch_number == len(train_loader):
|
||||||
|
processed = min(batch_number * args.batch_size, len(train_dataset))
|
||||||
|
elapsed = time.monotonic() - epoch_started_at
|
||||||
|
print(
|
||||||
|
f"[Эпоха {epoch}/{args.epochs}] train batch "
|
||||||
|
f"{batch_number}/{len(train_loader)}, "
|
||||||
|
f"страниц {processed}/{len(train_dataset)}, "
|
||||||
|
f"loss={running_loss / processed:.4f}, прошло {elapsed:.1f} с",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
model.eval()
|
||||||
|
correct = total = 0
|
||||||
|
print(
|
||||||
|
f"[Эпоха {epoch}/{args.epochs}] Validation, batches: {len(validation_loader)}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
with torch.inference_mode():
|
||||||
|
for batch_number, (images, targets, _) in enumerate(validation_loader, start=1):
|
||||||
|
images = images.to(device, non_blocking=device.type == "cuda")
|
||||||
|
targets = targets.to(device, non_blocking=device.type == "cuda")
|
||||||
|
if device.type == "cuda":
|
||||||
|
images = images.to(memory_format=torch.channels_last)
|
||||||
|
with torch.autocast(device_type=device.type, enabled=use_amp):
|
||||||
|
predicted = model(images).argmax(dim=1)
|
||||||
|
correct += int((predicted == targets).sum())
|
||||||
|
total += targets.size(0)
|
||||||
|
if batch_number % args.log_every == 0 or batch_number == len(validation_loader):
|
||||||
|
print(
|
||||||
|
f"[Эпоха {epoch}/{args.epochs}] validation batch "
|
||||||
|
f"{batch_number}/{len(validation_loader)}, "
|
||||||
|
f"текущая accuracy={correct / total:.4f}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
accuracy = correct / total
|
||||||
|
epoch_elapsed = time.monotonic() - epoch_started_at
|
||||||
|
remaining = epoch_elapsed * (args.epochs - epoch)
|
||||||
|
print(
|
||||||
|
f"Эпоха {epoch:02d}: loss={running_loss / len(train_dataset):.4f}, "
|
||||||
|
f"val_accuracy={accuracy:.4f}, время={epoch_elapsed:.1f} с, "
|
||||||
|
f"примерный остаток={remaining / 60:.1f} мин",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
if accuracy > best_accuracy:
|
||||||
|
best_accuracy = accuracy
|
||||||
|
epochs_without_improvement = 0
|
||||||
|
torch.save(
|
||||||
|
{
|
||||||
|
"state_dict": model.state_dict(),
|
||||||
|
"labels": labels,
|
||||||
|
"image_size": args.image_size,
|
||||||
|
"validation_accuracy": accuracy,
|
||||||
|
"architecture": "mobilenet_v3_small",
|
||||||
|
},
|
||||||
|
args.model,
|
||||||
|
)
|
||||||
|
print(f" Сохранена лучшая модель: {args.model}", flush=True)
|
||||||
|
else:
|
||||||
|
epochs_without_improvement += 1
|
||||||
|
if epochs_without_improvement >= args.patience:
|
||||||
|
print("Ранняя остановка: качество не улучшается.", flush=True)
|
||||||
|
break
|
||||||
|
print(
|
||||||
|
f"Готово за {(time.monotonic() - started_at) / 60:.1f} мин. "
|
||||||
|
f"Лучшая точность на страницах validation: {best_accuracy:.4f}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def predict(args: argparse.Namespace) -> None:
|
||||||
|
torch, _, _, _, _ = _dependencies()
|
||||||
|
device = _device(torch, args.device)
|
||||||
|
_configure_cuda(torch, device)
|
||||||
|
model, checkpoint = _load_checkpoint(args.model, device)
|
||||||
|
if device.type == "cuda":
|
||||||
|
model = model.to(memory_format=torch.channels_last)
|
||||||
|
files = _document_files(args.file) if args.file.is_dir() else [args.file]
|
||||||
|
if not files:
|
||||||
|
raise SystemExit("Поддерживаемые документы не найдены.")
|
||||||
|
for path in files:
|
||||||
|
try:
|
||||||
|
result = _predict(
|
||||||
|
model,
|
||||||
|
checkpoint["labels"],
|
||||||
|
checkpoint["image_size"],
|
||||||
|
path,
|
||||||
|
device,
|
||||||
|
args.threshold,
|
||||||
|
args.top_k,
|
||||||
|
)
|
||||||
|
print(json.dumps(asdict(result), ensure_ascii=False))
|
||||||
|
except Exception as exc:
|
||||||
|
print(json.dumps({"file": str(path), "error": str(exc)}, ensure_ascii=False))
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate(args: argparse.Namespace) -> None:
|
||||||
|
torch, _, _, _, _ = _dependencies()
|
||||||
|
device = _device(torch, args.device)
|
||||||
|
_configure_cuda(torch, device)
|
||||||
|
model, checkpoint = _load_checkpoint(args.model, device)
|
||||||
|
if device.type == "cuda":
|
||||||
|
model = model.to(memory_format=torch.channels_last)
|
||||||
|
classes = _scan_dataset(args.data)
|
||||||
|
expected_labels = checkpoint["labels"]
|
||||||
|
unknown_folders = sorted(set(classes) - set(expected_labels))
|
||||||
|
if unknown_folders:
|
||||||
|
raise SystemExit(f"Модель не знает типы: {', '.join(unknown_folders)}")
|
||||||
|
|
||||||
|
correct = accepted = total = 0
|
||||||
|
matrix: dict[str, Counter[str]] = {label: Counter() for label in classes}
|
||||||
|
for expected, paths in classes.items():
|
||||||
|
for path in paths:
|
||||||
|
result = _predict(
|
||||||
|
model,
|
||||||
|
expected_labels,
|
||||||
|
checkpoint["image_size"],
|
||||||
|
path,
|
||||||
|
device,
|
||||||
|
args.threshold,
|
||||||
|
top_k=1,
|
||||||
|
)
|
||||||
|
total += 1
|
||||||
|
accepted += int(result.accepted)
|
||||||
|
correct += int(result.document_type == expected)
|
||||||
|
matrix[expected][result.document_type] += 1
|
||||||
|
report = {
|
||||||
|
"documents": total,
|
||||||
|
"accuracy_with_unknown_as_error": round(correct / total, 6),
|
||||||
|
"accepted_share": round(accepted / total, 6),
|
||||||
|
"by_expected_type": {label: dict(counts) for label, counts in matrix.items()},
|
||||||
|
}
|
||||||
|
print(json.dumps(report, ensure_ascii=False, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser(description="Нейросетевой классификатор типов документов")
|
||||||
|
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||||
|
|
||||||
|
train_parser = subparsers.add_parser("train", help="Обучить модель")
|
||||||
|
train_parser.add_argument("--data", type=Path, required=True, help="Корень обучающего датасета")
|
||||||
|
train_parser.add_argument("--model", type=Path, required=True, help="Куда сохранить .pt")
|
||||||
|
train_parser.add_argument("--epochs", type=int, default=25)
|
||||||
|
train_parser.add_argument("--batch-size", type=int, default=DEFAULT_BATCH_SIZE)
|
||||||
|
train_parser.add_argument("--learning-rate", type=float, default=3e-4)
|
||||||
|
train_parser.add_argument("--validation-share", type=float, default=0.2)
|
||||||
|
train_parser.add_argument("--patience", type=int, default=5)
|
||||||
|
train_parser.add_argument("--image-size", type=int, default=DEFAULT_IMAGE_SIZE)
|
||||||
|
train_parser.add_argument(
|
||||||
|
"--workers",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_WORKERS,
|
||||||
|
help="Число процессов подготовки batches (для текущей реализации на Windows: 0)",
|
||||||
|
)
|
||||||
|
train_parser.add_argument(
|
||||||
|
"--log-every",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_LOG_EVERY,
|
||||||
|
help="Печатать прогресс каждые N batches",
|
||||||
|
)
|
||||||
|
train_parser.add_argument(
|
||||||
|
"--dataset-log-every",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_DATASET_LOG_EVERY,
|
||||||
|
help="Печатать прогресс чтения каждые N документов",
|
||||||
|
)
|
||||||
|
train_parser.add_argument("--seed", type=int, default=62625)
|
||||||
|
train_parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
||||||
|
train_parser.add_argument(
|
||||||
|
"--no-pretrained",
|
||||||
|
action="store_true",
|
||||||
|
help="Не скачивать ImageNet-веса (качество обычно ниже)",
|
||||||
|
)
|
||||||
|
train_parser.add_argument(
|
||||||
|
"--no-amp",
|
||||||
|
action="store_true",
|
||||||
|
help="Отключить mixed precision (для диагностики проблем CUDA)",
|
||||||
|
)
|
||||||
|
train_parser.set_defaults(handler=train)
|
||||||
|
|
||||||
|
predict_parser = subparsers.add_parser("predict", help="Определить тип файла или каталога")
|
||||||
|
predict_parser.add_argument("--model", type=Path, required=True)
|
||||||
|
predict_parser.add_argument("--file", type=Path, required=True)
|
||||||
|
predict_parser.add_argument("--threshold", type=float, default=DEFAULT_CONFIDENCE_THRESHOLD)
|
||||||
|
predict_parser.add_argument("--top-k", type=int, default=3)
|
||||||
|
predict_parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
||||||
|
predict_parser.set_defaults(handler=predict)
|
||||||
|
|
||||||
|
evaluate_parser = subparsers.add_parser("evaluate", help="Проверить модель на отдельном датасете")
|
||||||
|
evaluate_parser.add_argument("--model", type=Path, required=True)
|
||||||
|
evaluate_parser.add_argument("--data", type=Path, required=True)
|
||||||
|
evaluate_parser.add_argument("--threshold", type=float, default=DEFAULT_CONFIDENCE_THRESHOLD)
|
||||||
|
evaluate_parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
||||||
|
evaluate_parser.set_defaults(handler=evaluate)
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Iterable[str] | None = None) -> None:
|
||||||
|
parser = build_parser()
|
||||||
|
args = parser.parse_args(argv)
|
||||||
|
if hasattr(args, "threshold") and not 0.0 <= args.threshold <= 1.0:
|
||||||
|
parser.error("--threshold должен быть от 0 до 1")
|
||||||
|
if hasattr(args, "validation_share") and not 0.0 < args.validation_share < 1.0:
|
||||||
|
parser.error("--validation-share должен быть от 0 до 1")
|
||||||
|
if hasattr(args, "log_every") and args.log_every < 1:
|
||||||
|
parser.error("--log-every должен быть не меньше 1")
|
||||||
|
if hasattr(args, "dataset_log_every") and args.dataset_log_every < 1:
|
||||||
|
parser.error("--dataset-log-every должен быть не меньше 1")
|
||||||
|
args.handler(args)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
98
download_files.py
Normal file
98
download_files.py
Normal file
@@ -0,0 +1,98 @@
|
|||||||
|
"""Открыть в браузере ссылки из TXT для скачивания через браузерную сессию."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import time
|
||||||
|
import webbrowser
|
||||||
|
from pathlib import Path
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
|
||||||
|
def read_urls(txt_file: Path) -> list[str]:
|
||||||
|
"""Прочитать непустые HTTP/HTTPS-ссылки, кроме комментариев с #."""
|
||||||
|
urls: list[str] = []
|
||||||
|
with txt_file.open("r", encoding="utf-8-sig") as file:
|
||||||
|
for line_number, line in enumerate(file, start=1):
|
||||||
|
value = line.strip()
|
||||||
|
if not value or value.startswith("#"):
|
||||||
|
continue
|
||||||
|
if urlparse(value).scheme not in {"http", "https"}:
|
||||||
|
print(f"Строка {line_number} пропущена: это не HTTP/HTTPS-ссылка")
|
||||||
|
continue
|
||||||
|
urls.append(value)
|
||||||
|
return urls
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Последовательно открыть ссылки из TXT в системном браузере"
|
||||||
|
)
|
||||||
|
parser.add_argument("txt_file", type=Path, help="TXT-файл: одна ссылка на строку")
|
||||||
|
parser.add_argument(
|
||||||
|
"--delay",
|
||||||
|
type=float,
|
||||||
|
default=0.2,
|
||||||
|
help="Пауза между ссылками в секундах (по умолчанию 1.5)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--start",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Начать с указанного номера строки списка (по умолчанию 1)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--limit",
|
||||||
|
type=int,
|
||||||
|
help="Открыть не больше указанного числа ссылок",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if not args.txt_file.is_file():
|
||||||
|
parser.error(f"файл не найден: {args.txt_file}")
|
||||||
|
if args.delay < 0:
|
||||||
|
parser.error("--delay не может быть отрицательным")
|
||||||
|
if args.start < 1:
|
||||||
|
parser.error("--start должен быть не меньше 1")
|
||||||
|
if args.limit is not None and args.limit < 1:
|
||||||
|
parser.error("--limit должен быть не меньше 1")
|
||||||
|
|
||||||
|
all_urls = read_urls(args.txt_file)
|
||||||
|
urls = all_urls[args.start - 1 :]
|
||||||
|
if args.limit is not None:
|
||||||
|
urls = urls[: args.limit]
|
||||||
|
if not urls:
|
||||||
|
raise SystemExit("Нет ссылок для открытия.")
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Будет открыто ссылок: {len(urls)} из {len(all_urls)}. "
|
||||||
|
f"Начало с номера: {args.start}."
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
"Браузер сохраняет оригинальные имена и расширения. "
|
||||||
|
"Не закрывайте его до завершения."
|
||||||
|
)
|
||||||
|
|
||||||
|
failed = 0
|
||||||
|
for offset, url in enumerate(urls):
|
||||||
|
number = args.start + offset
|
||||||
|
try:
|
||||||
|
# new=0 просит браузер использовать существующее окно. Для download-
|
||||||
|
# ответа браузер обычно начинает загрузку, не оставляя новую вкладку.
|
||||||
|
opened = webbrowser.open(url, new=0, autoraise=False)
|
||||||
|
if not opened:
|
||||||
|
raise RuntimeError("системный браузер не принял ссылку")
|
||||||
|
print(f"[{offset + 1}/{len(urls)}] Открыта ссылка №{number}")
|
||||||
|
except Exception as error:
|
||||||
|
failed += 1
|
||||||
|
print(f"[{offset + 1}/{len(urls)}] ОШИБКА №{number}: {error}")
|
||||||
|
if offset + 1 < len(urls):
|
||||||
|
time.sleep(args.delay)
|
||||||
|
|
||||||
|
print(f"Готово. Передано браузеру: {len(urls) - failed}, ошибок: {failed}.")
|
||||||
|
if failed:
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
6
requirements.txt
Normal file
6
requirements.txt
Normal file
@@ -0,0 +1,6 @@
|
|||||||
|
--extra-index-url https://download.pytorch.org/whl/cu128
|
||||||
|
|
||||||
|
torch==2.11.0+cu128
|
||||||
|
torchvision==0.26.0+cu128
|
||||||
|
pillow==12.2.0
|
||||||
|
pypdfium2==5.12.1
|
||||||
Reference in New Issue
Block a user