Метод load_state_dict
Метод load_state_dict загружает состояние модели из словаря, полученного ранее с помощью метода state_dict. Первым параметром метод принимает словарь с весами, вторым параметром можно передать флаг строгой загрузки strict. Если strict=True (по умолчанию), то ключи загружаемого словаря должны точно совпадать с ключами текущей модели. Если некоторые ключи отсутствуют или лишние, будет вызвана ошибка.
Синтаксис
model.load_state_dict(state_dict, strict=True)
Пример
Давайте создадим простую модель, сохраним её состояние, а затем загрузим его в новую модель:
import torch
from torch import nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
def forward(self, x):
return self.fc(x)
model1 = SimpleModel()
model2 = SimpleModel()
state = model1.state_dict()
model2.load_state_dict(state)
Результат выполнения кода:
<All keys matched successfully>
Пример
Давайте загрузим состояние с проверкой ключей. Если структура модели не совпадает, будет ошибка:
import torch
from torch import nn
class ModelA(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
def forward(self, x):
return self.fc(x)
class ModelB(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(5, 3)
def forward(self, x):
return self.fc1(x)
model_a = ModelA()
model_b = ModelB()
state_a = model_a.state_dict()
model_b.load_state_dict(state_a, strict=True)
Результат выполнения кода:
RuntimeError: Error(s) in loading state_dict for ModelB:
Missing key(s) in state_dict: "fc1.weight", "fc1.bias"
Unexpected key(s) in state_dict: "fc.weight", "fc.bias"
Пример
Давайте загрузим состояние с отключённой строгой проверкой. Лишние или отсутствующие ключи будут проигнорированы:
import torch
from torch import nn
class ModelA(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
self.extra = nn.Linear(3, 1)
def forward(self, x):
return self.fc(x)
class ModelB(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
def forward(self, x):
return self.fc(x)
model_a = ModelA()
model_b = ModelB()
state_a = model_a.state_dict()
model_b.load_state_dict(state_a, strict=False)
Результат выполнения кода:
<All keys matched successfully>
Пример
Давайте загрузим предварительно сохранённое состояние из файла. Обычно модель сохраняется с помощью метода save:
import torch
from torch import nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
state_dict = torch.load('model.pt')
model.load_state_dict(state_dict)
Смотрите также
-
метод
state_dict,
который возвращает словарь с состоянием модели -
метод
parameters,
который возвращает итератор по обучаемым параметрам модели -
метод
train,
который переводит модель в режим обучения -
метод
eval,
который переводит модель в режим оценки