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

Метод register_module

Метод register_module класса Module позволяет динамически регистрировать вложенные модули в текущем модуле. Первым параметром метод принимает строку с именем субмодуля, вторым параметром - сам модуль, который нужно зарегистрировать. После регистрации модуль становится доступен как атрибут родительского модуля и учитывается в методах parameters, modules и других.

Синтаксис

module.register_module(name, submodule)

Пример

Давайте зарегистрируем линейный слой в пустом модуле:

import torch import torch.nn as nn model = nn.Module() model.register_module('fc', nn.Linear(10, 5)) print(model.fc)

Результат выполнения кода:

Linear(in_features=10, out_features=5, bias=True)

Пример

Зарегистрируем несколько модулей с разными именами:

import torch import torch.nn as nn model = nn.Module() model.register_module('layer1', nn.Linear(10, 20)) model.register_module('layer2', nn.Linear(20, 30)) model.register_module('layer3', nn.Linear(30, 5)) for name, param in model.named_parameters(): print(f"{name}: {param.shape}")

Результат выполнения кода:

layer1.weight: torch.Size([20, 10]) layer1.bias: torch.Size([20]) layer2.weight: torch.Size([30, 20]) layer2.bias: torch.Size([30]) layer3.weight: torch.Size([5, 30]) layer3.bias: torch.Size([5])

Пример

Динамически заменим зарегистрированный модуль:

import torch import torch.nn as nn model = nn.Module() model.register_module('fc', nn.Linear(10, 5)) print(f"Before: {model.fc}") model.register_module('fc', nn.Linear(20, 10)) print(f"After: {model.fc}")

Результат выполнения кода:

Before: Linear(in_features=10, out_features=5, bias=True) After: Linear(in_features=20, out_features=10, bias=True)

Пример

Используем register_module в пользовательском классе:

import torch import torch.nn as nn class MyModel(nn.Module): def __init__(self): super().__init__() self.register_module('fc1', nn.Linear(10, 20)) self.register_module('fc2', nn.Linear(20, 5)) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) model = MyModel() x = torch.randn(3, 10) res = model(x) print(res.shape)

Результат выполнения кода:

torch.Size([3, 5])

Смотрите также

  • метод add_module,
    который также добавляет вложенный модуль
  • метод register_parameter,
    который регистрирует параметр модуля
  • метод register_buffer,
    который регистрирует буфер модуля
  • метод get_submodule,
    который получает вложенный модуль по имени
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить