Класс Parameter
Класс Parameter является специальным типом тензора, который
автоматически регистрируется как параметр модуля при присвоении
атрибуту класса nn.Module. Основное отличие от обычного
тензора заключается в том, что параметры автоматически добавляются
в список параметров модуля и участвуют в процессе обучения.
При создании параметра первым аргументом передается тензор данных,
вторым аргументом можно указать необходимость вычисления градиента
(по умолчанию True).
Синтаксис
torch.nn.Parameter(data, requires_grad=True)
Пример
Давайте создадим простой параметр из тензора:
import torch
t = torch.tensor([1., 2., 3.])
param = torch.nn.Parameter(t)
print(param)
Результат выполнения кода:
Parameter containing:
tensor([1., 2., 3.], requires_grad=True)
Пример
Создадим параметр с отключенным вычислением градиента:
import torch
t = torch.tensor([1., 2., 3.])
param = torch.nn.Parameter(t, requires_grad=False)
print(param)
Результат выполнения кода:
Parameter containing:
tensor([1., 2., 3.], requires_grad=False)
Пример
Использование параметра в пользовательском модуле:
import torch
class MyLinear(torch.nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.weight = torch.nn.Parameter(
torch.randn(in_features, out_features)
)
self.bias = torch.nn.Parameter(
torch.zeros(out_features)
)
def forward(self, x):
return x @ self.weight + self.bias
model = MyLinear(3, 2)
for name, param in model.named_parameters():
print(f"{name}: {param.shape}")
Результат выполнения кода:
weight: torch.Size([3, 2])
bias: torch.Size([2])
Пример
Обновление параметра через оптимизатор:
import torch
torch.manual_seed(0)
param = torch.nn.Parameter(torch.tensor([5.]))
optimizer = torch.optim.SGD([param], lr=0.1)
loss = (param - 3) ** 2
loss.backward()
optimizer.step()
print(param)
Результат выполнения кода:
Parameter containing:
tensor([4.6000], requires_grad=True)
Пример
Преобразование обычного тензора в параметр:
import torch
t = torch.ones(3, 3)
param = torch.nn.Parameter(t)
print(param)
Результат выполнения кода:
Parameter containing:
tensor([
[1., 1., 1.],
[1., 1., 1.],
[1., 1., 1.],
], requires_grad=True)