Курс 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. Вакансии в Nebius
  3. Названия переменных
  4. Проверка запуска скрипта или импорта модуля
  5. Функция с **kwargs в Python
  6. Работа с collections в Python.
  7. Удаление дубликатов из списка с помощью dict.fromkeys
  8. Настройка логгера Logzero
  9. Получение текущей директории
  10. Роль object и type в Python
  11. Структурирование именованных констант
  12. Удаление первого элемента списка
  13. Разделение строки на пары ключ-значение.
  14. Регистрация на TenChat
  15. Разделение строк в Python
  16. Применение функции к списку
  17. Библиотека Emoji: использование смайлов в Python
  18. Хэш-функции в Python
  19. Подписка на Kaspersky Team
  20. Работа с файлами в Python
  21. Удаление специальных символов с помощью re.sub
  22. Объединение словарей в Python
  23. Обработка ошибок в JSON данных
  24. Генератор бросков кубиков
  25. Профилирование кода на Python
  26. Закрытие файла в Python
  27. Antigravity модуль
  28. Проверка на палиндром
  29. Генератор списка с условием if
  30. Встроенные функции Python
  31. Основы Python
  32. Оператор * в Python
  33. Работа с deque в Python
  34. Поиск наиболее частого элемента списке
  35. Сохранение и загрузка модели в PyTorch
  36. Метод __iand__ для пользовательских классов
  37. Форматирование строк в Python
  38. Удаление знаков препинания в Python
  39. Метод __ilshift__ для битового сдвига влево
  40. Улучшение читаемости кода в Python
  41. Python: библиотеки и функции
  42. Частичное применение функций в Python
  43. Списки в Python: основы
  44. Аннотации типов в Python
  45. Уникальные значения из списка

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