Курс Python → Тестирование модели в PyTorch

Для того чтобы эффективно оценивать работу нашей модели машинного обучения, необходимо определить метод тестирования. Этот метод позволит нам проверить качество работы модели на тестовом наборе данных и вывести точность предсказаний. Основное отличие метода тестирования от обучения заключается в том, что в процессе тестирования мы используем функцию model.eval(), чтобы перевести модель в режим тестирования. Также важно использовать torch.no_grad(), чтобы отключить вычисление градиента, поскольку во время тестирования обратное распространение не требуется.

Для начала необходимо перевести модель в режим тестирования с помощью функции model.eval(). Это гарантирует, что все слои модели будут работать в режиме тестирования, что может влиять на поведение некоторых слоев, таких как Dropout или BatchNorm. Затем мы используем torch.no_grad(), чтобы временно отключить автоматическое дифференцирование и вычисление градиента. Это позволяет ускорить процесс тестирования, поскольку не нужно хранить градиенты для обновления весов модели.


model.eval()

with torch.no_grad():
    for inputs, labels in test_loader:
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        test_loss += loss.item()
        _, predicted = torch.max(outputs, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

test_accuracy = correct / total

Наконец, после прохождения всех тестовых данных, мы вычисляем средние потери для всего тестового набора и общую точность предсказаний. Это позволяет оценить, насколько хорошо модель обучилась и способна предсказывать значения на новых данных. Результаты тестирования помогут нам понять, какие улучшения можно внести в модель для повышения ее эффективности и точности предсказаний.

Твои коллеги будут рады, поделись в

Автор урока

Дмитрий Комаровский
Дмитрий Комаровский

Автоматизация процессов
в КраснодарБанки.ру

Другие уроки курса "Python"

  1. Получение текущего времени в Python
  2. Проверка памяти объекта
  3. Встроенные функции Python
  4. Повторение элементов в Python
  5. Работа с NumPy.linalg
  6. Метод is_absolute() для PurePath
  7. Получение комбинаций в Python
  8. Уникальные значения из списка
  9. Закрытие файла в Python
  10. Python Аргументы по умолчанию
  11. Изменение IP-адреса в Python
  12. Основные функции и модули Python
  13. Переворот последовательности
  14. Генераторы в Python
  15. inspect в Python: анализ кода
  16. Работа с модулем random
  17. Переопределение метода len
  18. Динамическая типизация в Python
  19. Форматирование строк в Python
  20. Оператор «not» в Python
  21. Создание виртуальной среды
  22. Подсчет частоты элементов с Counter
  23. Метод hash в Python
  24. Операции с матрицами в Python
  25. Класс Counter() для подсчета элементов
  26. Атрибуты массивов в Numpy
  27. Добавление вложенных списков
  28. Объединение множеств в Python
  29. Удаление дубликатов из списка с помощью dict.fromkeys
  30. Работа со временем в Python
  31. Создание спинбокса в tkinter
  32. Операторы сравнения в Python
  33. Обработка ошибок в JSON данных
  34. Руководство по библиотеке pydantic
  35. Установка Python — Простое руководство
  36. Фильтры Pillow: NEAREST, BILINEAR, BICUBIC
  37. Вычисление логарифмов в Python
  38. Декораторы в Python
  39. Генераторы списков в Python
  40. Замена символов в Python
  41. Удаление знаков препинания в Python
  42. Работа с collections.Counter
  43. Запуск внешнего кода в Jupyter
  44. Метод rlshift для битового сдвига
  45. Объединение словарей в Python 3.5+
  46. Метод Event.wait() в Python

Marketello читают маркетологи из крутых компаний