Метод state_dict
Метод state_dict класса Module возвращает словарь Python,
содержащий все обучаемые параметры (веса и смещения) и постоянные буферы
модели. Ключами словаря выступают строковые имена параметров, значениями -
тензоры. Метод не принимает параметров.
Синтаксис
model.state_dict()
Пример
Создадим простую линейную модель и посмотрим её параметры:
import torch
import torch.nn as nn
model = nn.Linear(3, 2)
state = model.state_dict()
for key, value in state.items():
print(key, value.shape)
Результат выполнения кода:
weight torch.Size([2, 3])
bias torch.Size([2])
Как видно из результата, модель содержит два параметра - матрицу весов
размером 2x3 и вектор смещения размером 2.
Пример
Создадим более сложную модель с несколькими слоями:
import torch
import torch.nn as nn
class SimpleNN(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(10, 5)
self.fc2 = nn.Linear(5, 1)
self.relu = nn.ReLU()
def forward(self, x):
x = self.relu(self.fc1(x))
return self.fc2(x)
model = SimpleNN()
state = model.state_dict()
for key, value in state.items():
print(key, value.shape)
Результат выполнения кода:
fc1.weight torch.Size([5, 10])
fc1.bias torch.Size([5])
fc2.weight torch.Size([1, 5])
fc2.bias torch.Size([1])
Пример
Сохраним параметры модели в файл и восстановим их в другой модели:
import torch
import torch.nn as nn
model1 = nn.Linear(5, 3)
model2 = nn.Linear(5, 3)
torch.save(model1.state_dict(), 'model.pt')
model2.load_state_dict(torch.load('model.pt'))
print(torch.equal(model1.weight, model2.weight))
Результат выполнения кода:
True
В этом примере мы сохранили состояние первой модели в файл, а затем загрузили его во вторую модель. Проверка показала, что параметры полностью совпадают.
Смотрите также
-
метод
load_state_dict,
который загружает параметры модели из словаря -
метод
parameters,
который возвращает итератор по параметрам модели -
метод
named_parameters,
который возвращает итератор с именами и параметрами модели -
метод
buffers,
который возвращает итератор по постоянным буферам модели