Курс Python → Сохранение и загрузка модели в PyTorch

Для сохранения и загрузки модели в PyTorch необходимо использовать методы torch.save() и torch.load(). Для сохранения модели передайте model.state_dict() в качестве первого аргумента, это просто словарь, который содержит информацию о слоях модели и их параметрах (веса и смещения). Вторым аргументом укажите имя файла, в котором будет сохранена модель. Хорошей практикой является использование расширений .pth или .pt для сохранения моделей PyTorch. Также можно указать полный путь к файлу, если вы хотите сохранить модель в определенном каталоге.

Пример сохранения модели:


torch.save(model.state_dict(), "cifar_fc.pth")

Чтобы загрузить сохраненную модель для дальнейшего использования или логического вывода, используйте метод torch.load(). Затем можно загрузить параметры модели с помощью метода load_state_dict(). Это позволит восстановить состояние модели с сохраненными параметрами и продолжить обучение или использование модели для вывода.

Пример загрузки модели:


model = YourModelClass()
model.load_state_dict(torch.load("cifar_fc.pth"))
model.eval()

При загрузке модели убедитесь, что класс модели, для которой загружаются параметры, совпадает с классом модели, которая была сохранена. В противном случае возможны ошибки при загрузке параметров. Также рекомендуется использовать метод model.eval() после загрузки модели, чтобы переключить ее в режим оценки и отключить дополнительные режимы, такие как режим обучения.

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

Автор урока

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

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

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

  1. Реверс строки в Python
  2. Присоединение элементов коллекции
  3. Глобальные переменные в Python
  4. Сохранение и загрузка модели в PyTorch
  5. Метод get() для словарей
  6. Замер времени выполнения кода
  7. Извлечение новостей с newspaper3k
  8. Работа с итераторами в Python
  9. discard() — удаление элемента из множества
  10. Оператор is в Python
  11. Принципы Zen Python
  12. Резервирование символов в Python
  13. Метод matmul для умножения матриц
  14. Сортировка списка по индексам
  15. Методы сравнения множеств
  16. Генератор списка в Python
  17. Аннотации типов в Python
  18. Особенности запятых в Python
  19. Оператор «not» в Python
  20. Создание таблиц в Python с PrettyTable
  21. Управление ресурсами в Python
  22. Использование super() в Python
  23. Работа с набором данных CIFAR10 в PyTorch
  24. Работа с итераторами через срезы
  25. Метод __index__ в Python
  26. Оптимизация памяти с __slots__
  27. Управление виртуальными средами в Python
  28. Конвертация текстовых чисел с помощью Numerizer
  29. Классы данных в Python
  30. Магические методы в Python
  31. Списковое включение в Python
  32. Многострочные комментарии в Python
  33. Делегирование в Python
  34. Хранение данных с помощью dataclasses
  35. Codecademy в Telegram
  36. Оценка выражений генератора в Python
  37. Функции map, filter и reduce
  38. Область видимости переменных в Python
  39. Метод rlshift для битового сдвига
  40. Копирование файлов с shutil()
  41. Удаление элементов из списка в Python
  42. Структура данных словарь в Python
  43. Кортеж в Python: создание, доступ, изменение
  44. Повторение и перенос строки
  45. Импорт модулей и пакетов в Python

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