Ответ
Для классификации простых бинарных фигур (масок) я бы использовал сверточную нейронную сеть (CNN), но начал бы с очень простой архитектуры, так как задача не требует глубоких признаков.
Архитектура модели (PyTorch):
import torch.nn as nn
class ShapeClassifier(nn.Module):
def __init__(self, num_classes=3):
super().__init__()
# Вход: 1 канал (бинарная маска)
self.features = nn.Sequential(
nn.Conv2d(1, 8, kernel_size=3, padding=1), # 8 фильтров
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(8, 16, kernel_size=3, padding=1), # 16 фильтров
nn.ReLU(),
nn.MaxPool2d(2),
nn.Flatten()
)
# Размер входа в полносвязный слой зависит от размера изображения.
# Для изображения 64x64 после двух пулингов (64->32->16) размер будет: 16 * 16 * 16 = 4096
self.classifier = nn.Linear(4096, num_classes)
def forward(self, x):
x = self.features(x)
x = self.classifier(x)
return x
Ключевые этапы:
- Подготовка данных: Генерация синтетического датасета с бинарными масками фигур (1 - фигура, 0 - фон). Добавляю аугментации: небольшие повороты, сдвиги, шум, чтобы модель обобщала.
- Обучение: Использую
CrossEntropyLossи оптимизатор Adam. Так как фигуры простые, обучение сходится быстро даже на небольшой выборке (несколько тысяч примеров). - Почему простая CNN? Глубокие предобученные сети (ResNet) избыточны для бинарных масок. Моя легкая сеть быстро обучается и меньше склонна к переобучению на таком синтетическом датасете.
- Альтернативный подход: Для абсолютной интерпретируемости можно было бы использовать классические методы компьютерного зрения (например, поиск контуров и анализ числа вершин с помощью
cv2.approxPolyDP), но нейросетевое решение более устойчиво к артефактам рендеринга и шуму.