commit 767a4c0a93752dee6febf8fe70c634589c27d509 Author: e.vlasov Date: Mon Jul 27 10:16:45 2026 +0300 refs #62625 добавил прототип классификатора типов документов diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..95057df --- /dev/null +++ b/.gitignore @@ -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 diff --git a/README.md b/README.md new file mode 100644 index 0000000..0d5bfb0 --- /dev/null +++ b/README.md @@ -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`. Скрипт не +переименовывает файлы: браузер сохраняет оригинальное имя из ответа сервера. diff --git a/document_type_classifier.py b/document_type_classifier.py new file mode 100644 index 0000000..b59d248 --- /dev/null +++ b/document_type_classifier.py @@ -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() diff --git a/download_files.py b/download_files.py new file mode 100644 index 0000000..5dc8fee --- /dev/null +++ b/download_files.py @@ -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() diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..a571328 --- /dev/null +++ b/requirements.txt @@ -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