Курс 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. Функция findall() для поиска вхождений строки
  2. Работа с JSON данными в Python
  3. Python reversed() функция
  4. Декораторы в Python
  5. Методы работы со строками в Python
  6. Генераторные функции в Python
  7. Копирование списков в Python
  8. Обработка ошибок в Python
  9. Многопоточность и асинхронное программирование в Python
  10. Объединение списков с помощью zip
  11. Асинхронное программирование с asyncio
  12. Установка Home Assistant
  13. Работа со строками в Python.
  14. Функции с дополнением
  15. Метод rlshift для битового сдвига
  16. Исключение NotImplementedError
  17. Операция += для списков
  18. Объединение строк с помощью метода join
  19. Очистка данных с помощью pandas
  20. Структурирование именованных констант
  21. Работа с timedelta в Python
  22. Объединение объектов в Python
  23. Работа с датой и временем в Python
  24. Вывод сложных структур данных с помощью pprint
  25. Обработка исключений в Python
  26. Разделение строки в Python
  27. Блок else в циклах Python
  28. Сериализация данных в JSON с помощью json.dumps
  29. Замена текста с помощью sub
  30. Пропуск строк в файле с itertools
  31. Установка библиотек в Python
  32. Замена элементов в списке с помощью генераторов списков
  33. Логирование с Logzero
  34. Создание копии списка в Python
  35. Кортеж в Python: создание, доступ, изменение
  36. Генераторы в Python
  37. Непрерывная проверка в Python
  38. Подсказки при вводе данных в Python
  39. Принципы Zen of Python
  40. Удаление ресурса в Python
  41. Принцип одной функции
  42. Оформление кода на Python
  43. Фильтрация списка чисел
  44. Очистка входных данных
  45. Оператор обр. импликации
  46. Сортировка данных в Python

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