Загрузка весов в PyTorch
Если нужно вернуть сохранённые
числа, файл читают функцией
torch.load. Из файла
берут только веса, а не произвольный
код. Затем метод load_state_dict
у нового модуля той же архитектуры
подставляет тензоры на место
параметров:
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)
torch.manual_seed(0)
net = TinyNet()
path = "weights.pt"
torch.save(net.state_dict(), path)
fresh = TinyNet()
before = fresh.layer.weight[0, 0].item()
stored = torch.load(path, weights_only=True)
fresh.load_state_dict(stored)
after = fresh.layer.weight[0, 0].item()
print(before, after)
# выведет 0.1871093362569809 -0.005293981172144413
После подстановки веса совпадают с исходной моделью. Сравним матрицу линейного слоя у сохранённой и загруженной сети:
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)
torch.manual_seed(0)
net = TinyNet()
path = "weights.pt"
torch.save(net.state_dict(), path)
fresh = TinyNet()
stored = torch.load(path, weights_only=True)
fresh.load_state_dict(stored)
same = torch.allclose(
fresh.layer.weight, net.layer.weight
)
print(same) # выведет True
Создайте модуль с линейным слоем
1 на 1, запишите его
параметры в one.pt.
Соберите второй такой же модуль,
прочитайте файл и подставьте
данные в параметры. Выведите
один элемент матрицы весов до
подстановки и после неё.
Опишите сеть с линейным
преобразованием 2 на 2,
сохраните веса в pair.pt.
На новом объекте той же формы
восстановите параметры из файла
и выведите, совпадают ли
матрицы весов.
Соберите блок 3 на 1,
положите словарь параметров
в bias.pt. Создайте
ещё один блок той же схемы,
загрузите данные с диска
и выведите форму вектора
смещения после восстановления.