Метод 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,
который получает вложенный модуль по имени