РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
473 of 769 menu

Метод 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,
    который возвращает все параметры модели
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить