Метод register_parameter
Метод register_parameter класса Module позволяет
зарегистрировать обучаемый параметр в модели. Первым параметром
метод принимает имя параметра в виде строки, вторым - сам параметр
в виде тензора Parameter. Зарегистрированный параметр
автоматически добавляется в список параметров модели и будет
обучаться в процессе оптимизации.
Синтаксис
module.register_parameter(name, param)
Пример
Давайте создадим простой модуль и зарегистрируем в нём один параметр:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_parameter('weight', nn.Parameter(torch.tensor([1.0, 2.0, 3.0])))
model = MyModule()
print(model.weight)
Результат выполнения кода:
Parameter containing:
tensor([1., 2., 3.], requires_grad=True)
Пример
Зарегистрируем несколько параметров с разными размерами:
import torch
import torch.nn as nn
class LinearLayer(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
weight = nn.Parameter(torch.randn(out_features, in_features))
bias = nn.Parameter(torch.randn(out_features))
self.register_parameter('weight', weight)
self.register_parameter('bias', bias)
torch.manual_seed(0)
layer = LinearLayer(3, 2)
print(layer.weight)
print(layer.bias)
Результат выполнения кода:
Parameter containing:
tensor([[ 1.5410, -0.2934, -2.1788],
[ 0.5684, -1.0845, -1.3986]], requires_grad=True)
Parameter containing:
tensor([0.4033, 0.8380], requires_grad=True)
Пример
Параметр можно зарегистрировать как None, если он ещё не определён:
import torch
import torch.nn as nn
class FlexibleModule(nn.Module):
def __init__(self):
super().__init__()
self.register_parameter('param', None)
model = FlexibleModule()
print(model.param)
# Позже параметр можно установить
model.param = nn.Parameter(torch.tensor([5.0, 6.0]))
print(model.param)
Результат выполнения кода:
None
Parameter containing:
tensor([5., 6.], requires_grad=True)
Пример
Зарегистрированные параметры автоматически появляются в списке параметров модели:
import torch
import torch.nn as nn
class MyNetwork(nn.Module):
def __init__(self):
super().__init__()
self.register_parameter('weight', nn.Parameter(torch.ones(3)))
self.register_parameter('bias', nn.Parameter(torch.zeros(3)))
model = MyNetwork()
for name, param in model.named_parameters():
print(f'{name}: {param.shape}')
Результат выполнения кода:
weight: torch.Size([3])
bias: torch.Size([3])
Смотрите также
-
метод
register_buffer,
который регистрирует буфер (необучаемый тензор) -
метод
register_module,
который регистрирует вложенный подмодуль -
метод
named_parameters,
который возвращает имена и параметры модели -
метод
parameters,
который возвращает все параметры модели