Навчання моделі класифікації зображень (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 тижнів |
Хочете таку ж модель для свого проєкту? Зв'яжіться з нами — обговоримо ваш набір даних і підберемо оптимальне рішення. Замовте навчання моделі та отримайте консультацію з точною оцінкою термінів.







