Files
request-document-classifier/document_type_classifier.py

698 lines
29 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Обучение и запуск нейросетевого классификатора типов документов.
Поддерживаются 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()