Оценка качества модели

Введение в глубокое обучение с PyTorch

Jasmin Ludolf

Senior Data Science Content Developer, DataCamp

Обучение, валидация и тестирование

$$

  • Набор данных обычно делится на три подмножества:
Доля данных Роль
Обучающая 80–90% Настраивает параметры модели
Валидационная 10–20% Подбирает гиперпараметры
Тестовая 5–10% Оценивает итоговое качество модели

$$

  • Отслеживайте потери и точность на обучающей и валидационной выборках
Введение в глубокое обучение с PyTorch

Вычисление потерь на обучении

$$

На каждой эпохе:

  • Суммируйте потери по всем батчам в загрузчике данных
  • Вычислите среднее значение потерь по итогам эпохи
training_loss = 0.0

for inputs, labels in trainloader: # Run the forward pass outputs = model(inputs) # Compute the loss loss = criterion(outputs, labels)
# Backpropagation loss.backward() # Compute gradients optimizer.step() # Update weights optimizer.zero_grad() # Reset gradients
# Calculate and sum the loss training_loss += loss.item()
epoch_loss = training_loss / len(trainloader)
Введение в глубокое обучение с PyTorch

Вычисление потерь на валидации

validation_loss = 0.0
model.eval() # Put model in evaluation mode


with torch.no_grad(): # Disable gradients for efficiency
for inputs, labels in validationloader: # Run the forward pass outputs = model(inputs) # Calculate the loss loss = criterion(outputs, labels) validation_loss += loss.item() epoch_loss = validation_loss / len(validationloader) # Compute mean loss
model.train() # Switch back to training mode
Введение в глубокое обучение с PyTorch

Переобучение

пример переобучения

Введение в глубокое обучение с PyTorch

Вычисление точности с помощью torchmetrics

import torchmetrics


# Create accuracy metric metric = torchmetrics.Accuracy(task="multiclass", num_classes=3)
for features, labels in dataloader: outputs = model(features) # Forward pass # Compute batch accuracy (keeping argmax for one-hot labels) metric.update(outputs, labels.argmax(dim=-1))
# Compute accuracy over the whole epoch accuracy = metric.compute()
# Reset metric for the next epoch metric.reset()
Введение в глубокое обучение с PyTorch

Давайте потренируемся!

Введение в глубокое обучение с PyTorch

Preparing Video For Download...