РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
⊗pytoPmSvWr 93 of 95 menu
◀ ▶

Запись весов в PyTorch

Чтобы перенести веса между запусками, словарь параметров записывают в файл на диске. Функция torch.save принимает данные и путь. Сохраним только словарь весов, а не весь объект модели:

import os import torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.layer = nn.Linear(2, 3) def forward(self, x): return self.layer(x) net = TinyNet() path = "weights.pt" torch.save(net.state_dict(), path) print(os.path.isfile(path)) # выведет True

После записи файл можно найти по указанному имени. Проверим, что на диске появился именно выбранный путь:

import os import torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.layer = nn.Linear(2, 3) def forward(self, x): return self.layer(x) net = TinyNet() path = "weights.pt" torch.save(net.state_dict(), path) print(os.path.basename(path)) # выведет weights.pt

Соберите модуль с линейным слоем 1 на 1, снимите с него словарь параметров и запишите его в файл checkpoint.pt. Выведите, что файл существует на диске.

Опишите блок с линейным преобразованием 4 на 2. Сохраните его параметры в layer.pt и выведите имя файла без каталога.

Создайте сеть с одним линейным слоем на два входа и три выхода. Запишите словарь её весов в net.pt. Выведите, существует ли этот файл.

← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить