Класс ParameterDict
Класс ParameterDict из модуля torch.nn
представляет собой словарь, предназначенный для хранения
параметров модели в виде объектов Parameter.
Он наследуется от torch.nn.Module и позволяет
обращаться к параметрам по строковым ключам, как к
обычным атрибутам. Это особенно удобно, когда
количество параметров динамическое или их имена
формируются во время выполнения программы.
Синтаксис
torch.nn.ParameterDict(parameters=None)
Основные возможности
Класс ParameterDict поддерживает все операции
обычного словаря Python: добавление, удаление,
изменение элементов, а также итерацию по ключам и
значениям. При этом все содержащиеся в нём параметры
автоматически регистрируются как параметры модуля и
будут учитываться при вызове parameters.
Пример
Давайте создадим пустой словарь параметров и добавим в него два параметра с разными именами:
import torch
from torch.nn import ParameterDict, Parameter
params = ParameterDict()
params['weight'] = Parameter(torch.randn(3, 4))
params['bias'] = Parameter(torch.zeros(4))
for name, param in params.items():
print(f"{name}: {param.shape}")
Результат выполнения кода:
weight: torch.Size([3, 4])
bias: torch.Size([4])
Пример
Инициализируем словарь параметров сразу при создании, передав обычный словарь:
import torch
from torch.nn import ParameterDict, Parameter
torch.manual_seed(0)
params = ParameterDict({
'w1': Parameter(torch.ones(2, 3)),
'w2': Parameter(torch.randn(3, 2)),
})
for name, param in params.items():
print(f"{name}:\\n{param}\\n")
Результат выполнения кода:
w1:
tensor([[1., 1., 1.],
[1., 1., 1.]])
w2:
tensor([[ 1.5410, -0.2934],
[-2.1788, 0.5684],
[-1.0845, -1.3986]])
Пример
Продемонстрируем, как ParameterDict
автоматически регистрирует параметры. Доступ к ним
можно получить через метод parameters:
import torch
from torch.nn import ParameterDict, Parameter
params = ParameterDict()
params['weight'] = Parameter(torch.ones(2, 2))
params['bias'] = Parameter(torch.zeros(2))
for param in params.parameters():
print(param)
Результат выполнения кода:
tensor([[1., 1.],
[1., 1.]])
tensor([0., 0.])
Пример
Словарь параметров можно использовать внутри пользовательского модуля, чтобы динамически управлять параметрами:
import torch
from torch.nn import Module, ParameterDict, Parameter
class DynamicLinear(Module):
def __init__(self, in_features, out_features):
super().__init__()
self.params = ParameterDict({
'weight': Parameter(torch.randn(out_features, in_features)),
'bias': Parameter(torch.zeros(out_features)),
})
def forward(self, x):
return x @ self.params['weight'].T + self.params['bias']
torch.manual_seed(0)
model = DynamicLinear(3, 2)
x = torch.randn(1, 3)
res = model(x)
print(res)
Результат выполнения кода:
tensor([[-0.3252, 3.0241]])
Смотрите также
-
класс
ParameterList,
который хранит параметры в виде списка -
класс
ModuleDict,
который хранит модули в виде словаря -
класс
UninitializedParameter,
который представляет неинициализированный параметр -
класс
Buffer,
который представляет буфер, не являющийся параметром