Навчання моделі класифікації зображень (ResNet, EfficientNet, ViT)
Ми часто стикаємося з ситуацією: клієнт приносить набір даних з 500 зображень, розбитих на 10 класів, і просить навчити ResNet-50. Здавалося б, типова задача, але без правильної стратегії результат — 65% accuracy та перенавчання. Наш досвід показує, що успіх визначається трьома факторами: вибір архітектури під обсяг даних, боротьба з дисбалансом класів та підготовка моделі до production. Розповімо, як ми вирішуємо ці завдання і що входить у нашу роботу.
Наші показники: 10+ інженерів, 5+ років досвіду у Computer Vision, 50+ виконаних проєктів, компанія на ринку 5 років. Своєчасне виявлення дефектів за допомогою навчених моделей економить клієнтам до 2 млн грн на рік. Гарантуємо якість моделей — кожна проходить валідацію на тестовому наборі. Сертифіковані інженери з 5-річним досвідом у Computer Vision.
Кейс: класифікація дефектів лиття на виробництві
До нас звернувся завод з лиття алюмінієвих деталей. Завдання: класифікувати 12 типів дефектів за фото з конвеєра. Дані: 15 000 зображень, сильний дисбаланс — один клас становив 60% вибірки, рідкісні дефекти — менше 100 прикладів. Ми обрали EfficientNet-B4 з попереднім навчанням на ImageNet, застосували Focal Loss (gamma=2.0) та агресивну аугментацію (RandomErasing, CutMix). Macro F1 виріс з 0.32 до 0.89. Час інференсу на CPU з ONNX — 8 мс на зображення. Завод впровадив модель у лінію контролю якості, скоротивши витрати на ручну перевірку на 1.5 млн грн щорічно.
Вибір архітектури під набір даних
| Набір даних | Рекомендація | Чому |
|---|---|---|
| <500 прикладів/клас | EfficientNet-B0/B2 (frozen → partial unfreeze) | Менше параметрів, менше перенавчання |
| 500–5000/клас | EfficientNet-B4, ConvNeXt-T, ResNet-50 | Баланс точності та швидкості навчання |
| >5000/клас | ViT-B/16, ConvNeXt-S/B | Трансформери розкриваються на великих даних |
| Медицина, мало даних | ResNet-50 з pretrain на MedNet | Доменний pretrain важливіший за архітектуру |
| Edge / мобайл | MobileNetV3-Large, EfficientNet-Lite | Latency < 10ms на телефоні |
Покроковий процес навчання
- Аналіз набору даних — оцінка якості розмітки, підрахунок розподілу класів, виявлення шуму.
- Вибір стратегії — підбір архітектури, аугментацій та функції втрат (наприклад, Focal Loss для дисбалансу).
- Експерименти — запуск серії навчань з різними конфігами на Weights & Biases.
- Fine-tuning — двофазне навчання ViT: спочатку тільки head, потім поступове розморожування блоків.
- Тестування — оцінка на валідації з TTA, експорт в ONNX та перевірка latency.
- Передача моделі — документація, скрипти інференсу, контейнеризація (Docker) для простої інтеграції.
Двофазне навчання ViT
ViT на 300 прикладах без правильної стратегії дасть accuracy 65% там, де EfficientNet дасть 84%. Причина: attention heads на маленькому наборі даних перенавчаються швидше, ніж convolutional inductive bias в ResNet. Рішення — агресивна заморозка + поступове розморожування:
Код двофазного навчання ViT
import timm import torch import torch.nn as nn from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR def train_vit_two_phase( num_classes: int, train_loader, val_loader, phase1_epochs: int = 20, phase2_epochs: int = 60, device: str = 'cuda' ) -> nn.Module: model = timm.create_model( 'vit_base_patch16_224', pretrained=True, num_classes=num_classes, drop_rate=0.1, drop_path_rate=0.1 ).to(device) # ФАЗА 1: навчаємо тільки head for name, param in model.named_parameters(): param.requires_grad = 'head' in name opt1 = AdamW( filter(lambda p: p.requires_grad, model.parameters()), lr=1e-3, weight_decay=0.05 ) sched1 = CosineAnnealingLR(opt1, T_max=phase1_epochs, eta_min=1e-5) for epoch in range(phase1_epochs): _train_epoch(model, train_loader, opt1, device) sched1.step() # ФАЗА 2: розморожуємо останні 6 блоків (з 12) for name, param in model.named_parameters(): if any(f'blocks.{i}.' in name for i in range(6, 12)): param.requires_grad = True opt2 = AdamW([ {'params': model.head.parameters(), 'lr': 5e-5}, {'params': [p for n, p in model.named_parameters() if 'blocks' in n and p.requires_grad], 'lr': 5e-6} ], weight_decay=0.05) sched2 = CosineAnnealingLR(opt2, T_max=phase2_epochs, eta_min=1e-7) for epoch in range(phase2_epochs): _train_epoch(model, train_loader, opt2, device) _validate(model, val_loader, device) sched2.step() return model def _train_epoch(model, loader, optimizer, device): model.train() criterion = nn.CrossEntropyLoss(label_smoothing=0.1) for images, labels in loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() loss = criterion(model(images), labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() Боротьба з дисбалансом класів
Accuracy 94% при дисбалансі 1:200 — це зазвичай означає, що модель передбачає тільки мажоритарний клас. Метрики для дисбалансу: macro F1, balanced accuracy, per-class recall. Lin et al. запропонували Focal Loss, який ми адаптуємо під задачу. Порівняно з CrossEntropyLoss, Focal Loss підвищує macro F1 в 2–3 рази на сильному дисбалансі.
Код Focal Loss
import torch import torch.nn.functional as F class FocalLoss(nn.Module): def __init__(self, alpha: float = 1.0, gamma: float = 2.0, class_weights: torch.Tensor = None): super().__init__() self.alpha = alpha self.gamma = gamma self.class_weights = class_weights def forward(self, inputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: ce_loss = F.cross_entropy( inputs, targets, weight=self.class_weights, reduction='none' ) pt = torch.exp(-ce_loss) focal_loss = self.alpha * (1 - pt) ** self.gamma * ce_loss return focal_loss.mean() # Ваги класів обернено пропорційно частоті class_counts = torch.tensor([10000, 500, 80], dtype=torch.float) class_weights = (1.0 / class_counts).to(device) criterion = FocalLoss(gamma=2.0, class_weights=class_weights) Як прискорити інференс моделі?
Test Time Augmentation (TTA) — простий спосіб підняти accuracy на 1–3% порівняно з одиночною інференцією без перенавчання: прогоняємо кілька аугментованих версій зображення та усереднюємо передбачення.
import torchvision.transforms as T class TTAClassifier: def __init__(self, model: nn.Module, n_augments: int = 5): self.model = model.eval() self.n_augments = n_augments self.base_transform = T.Compose([ T.Resize(256), T.CenterCrop(224), T.ToTensor(), T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]) ]) self.tta_transforms = [ T.Compose([T.Resize(256), T.CenterCrop(224), T.RandomHorizontalFlip(p=1.0), T.ToTensor(), T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])]), T.Compose([T.Resize(224), T.ToTensor(), T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])]), ] @torch.no_grad() def predict(self, image) -> torch.Tensor: preds = [ torch.softmax(self.model(self.base_transform(image).unsqueeze(0)), dim=1) ] for transform in self.tta_transforms: preds.append( torch.softmax(self.model(transform(image).unsqueeze(0)), dim=1) ) return torch.stack(preds).mean(dim=0) Як експортувати модель для production?
# ONNX export — для CPU/TensorRT деплою dummy = torch.randn(1, 3, 224, 224, device='cpu') torch.onnx.export( model.cpu().eval(), dummy, 'classifier.onnx', opset_version=17, input_names=['image'], output_names=['logits'], dynamic_axes={'image': {0: 'batch'}, 'logits': {0: 'batch'}} ) Використання ONNX дозволяє знизити витрати на GPU до 40% за рахунок оптимізованого інференсу на CPU.
Що входить у роботу
- Аналіз набору даних: оцінка якості, дисбаланс, підбір аугментацій.
- Вибір архітектури та стратегії навчання під ваші дані.
- Реалізація пайплайну на PyTorch / Hugging Face.
- Підбір гіперпараметрів, логування метрик (Weights & Biases).
- Тестування з TTA та експорт в ONNX / TensorRT.
- Документація, файли моделі та скрипти інференсу.
- Консультація з інтеграції вашим розробникам.
Терміни
| Завдання | Термін |
|---|---|
| Fine-tuning готової архітектури (готові дані) | 1–3 тижні |
| Навчання з нуля + аугментації + оптимізація | 4–7 тижнів |
| Розробка кастомної архітектури під домен | 8–14 тижнів |
Хочете таку ж модель для свого проєкту? Зв'яжіться з нами — обговоримо ваш набір даних і підберемо оптимальне рішення. Замовте навчання моделі та отримайте консультацію з точною оцінкою термінів.







