Класс ParameterList
Класс ParameterList предназначен для хранения списка параметров модели в PyTorch.
Он наследуется от nn.Module и позволяет удобно управлять коллекцией параметров,
обеспечивая их правильную регистрацию в модели. Первым параметром конструктор принимает
итерируемый объект с параметрами, например, список тензоров или модулей.
Синтаксис
torch.nn.ParameterList(parameters=None)
Пример
Давайте создадим список параметров из нескольких тензоров:
import torch
import torch.nn as nn
params = nn.ParameterList([
nn.Parameter(torch.tensor([1.0, 2.0, 3.0])),
nn.Parameter(torch.tensor([4.0, 5.0, 6.0])),
nn.Parameter(torch.tensor([7.0, 8.0, 9.0])),
])
for i, p in enumerate(params):
print(f"Param {i}: {p}")
Результат выполнения кода:
Param 0: Parameter containing:
tensor([1., 2., 3.], requires_grad=True)
Param 1: Parameter containing:
tensor([4., 5., 6.], requires_grad=True)
Param 2: Parameter containing:
tensor([7., 8., 9.], requires_grad=True)
Пример
Давайте создадим список параметров с использованием модулей:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.params = nn.ParameterList([
nn.Parameter(torch.randn(2, 3)),
nn.Parameter(torch.randn(2, 3)),
])
def forward(self, x):
res = x
for p in self.params:
res = res @ p
return res
model = MyModule()
x = torch.randn(3, 2)
out = model(x)
print(out.shape)
Результат выполнения кода:
torch.Size([3, 3])
Пример
Давайте добавим новый параметр в существующий список:
import torch
import torch.nn as nn
params = nn.ParameterList([
nn.Parameter(torch.tensor([1.0, 2.0])),
nn.Parameter(torch.tensor([3.0, 4.0])),
])
print(f"До добавления: {len(params)} параметров")
params.append(nn.Parameter(torch.tensor([5.0, 6.0])))
print(f"После добавления: {len(params)} параметров")
print(params[-1])
Результат выполнения кода:
До добавления: 2 параметров
После добавления: 3 параметров
Parameter containing:
tensor([5., 6.], requires_grad=True)
Пример
Давайте изменим параметр в списке по индексу:
import torch
import torch.nn as nn
params = nn.ParameterList([
nn.Parameter(torch.tensor([1.0, 2.0])),
nn.Parameter(torch.tensor([3.0, 4.0])),
])
print(f"До изменения: {params[0]}")
params[0] = nn.Parameter(torch.tensor([10.0, 20.0]))
print(f"После изменения: {params[0]}")
Результат выполнения кода:
До изменения: Parameter containing:
tensor([1., 2.], requires_grad=True)
После изменения: Parameter containing:
tensor([10., 20.], requires_grad=True)
Пример
Давайте создадим модель с несколькими линейными слоями с общим списком параметров:
import torch
import torch.nn as nn
class MultiLinear(nn.Module):
def __init__(self, in_features, out_features, num_layers):
super().__init__()
self.weights = nn.ParameterList([
nn.Parameter(torch.randn(in_features, out_features))
for _ in range(num_layers)
])
def forward(self, x):
res = []
for w in self.weights:
res.append(x @ w)
return torch.stack(res)
torch.manual_seed(0)
model = MultiLinear(3, 2, 3)
x = torch.randn(2, 3)
out = model(x)
print(out.shape)
Результат выполнения кода:
torch.Size([3, 2, 2])
Смотрите также
-
класс
ParameterDict,
который представляет собой словарь параметров -
класс
ModuleDict,
который хранит словарь подмодулей -
класс
UninitializedParameter,
который представляет неинициализированный параметр -
класс
Buffer,
который используется для хранения буферов модели