Словарь весов модуля в PyTorch
Внутри модуля обучаемые числа
имеют имена: у линейного слоя
это матрица весов и вектор смещения.
Метод state_dict собирает их
в обычный словарь Python. Создадим
простую сеть и выведем ключи:
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()
weights = net.state_dict()
print(list(weights.keys()))
# выведет ['layer.weight', 'layer.bias']
Каждый ключ указывает на параметр внутри модели, а значение - тензор с числами. Такой словарь удобно передавать дальше, не таская за собой весь объект сети:
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()
weights = net.state_dict()
print(weights["layer.weight"].shape)
# выведет torch.Size([3, 2])
Соберите модуль с линейным слоем
на 2 входа и 4 выхода.
Получите словарь его параметров
и выведите список имён ключей.
Опишите блок с линейным преобразованием
3 на 1. Выведите форму
тензора, который лежит по ключу
с именем смещения слоя.
Создайте сеть с одним линейным
слоем 1 на 2.
Снимите с неё словарь параметров
и выведите, сколько ключей в нём
находится.