Запись весов в 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. Выведите,
существует ли этот файл.