Проверка имён при загрузке в PyTorch
При подстановке словаря весов
метод 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()
stored = net.state_dict()
fresh = TinyNet()
fresh.load_state_dict(stored)
same = torch.allclose(
fresh.layer.weight, net.layer.weight
)
print(same) # выведет True
Если в словаре появится ключ, которого нет у модели, подстановка прервётся ошибкой. Поймаем её и выведем тип исключения:
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()
stored = net.state_dict()
stored["ghost.weight"] = torch.randn(1)
try:
net.load_state_dict(stored)
except Exception as err:
print(type(err).__name__)
# выведет RuntimeError
Создайте модуль с линейным слоем
2 на 2, сохраните
словарь весов в переменную
и подставьте его во второй
такой же модуль. Выведите,
совпадают ли матрицы весов.
Соберите блок 1 на 3,
получите его словарь параметров,
добавьте в него посторонний
ключ с маленьким тензором
и попытайтесь подставить
данные в исходный модуль.
Поймайте сбой и выведите
имя класса исключения.
Опишите сеть с линейным
преобразованием 3 на 1,
запишите веса на диск
в tiny.pt. На новом
объекте той же схемы прочитайте
файл и восстановите параметры.
Выведите один элемент вектора
смещения после подстановки.