Класс ModuleDict
Класс ModuleDict из модуля torch.nn предназначен для хранения набора дочерних модулей (слоёв) в виде словаря. Он работает аналогично стандартному словарю Python, но при этом все добавленные модули автоматически регистрируются как параметры модели, что позволяет оптимизатору обновлять их веса. Это особенно полезно, когда количество или тип слоёв заранее неизвестны и определяются динамически, например, на основе входных данных или конфигурации. Класс наследуется от nn.Module и поддерживает большинство привычных методов словаря.
Синтаксис
class torch.nn.ModuleDict(modules=None)
Параметры:
-
modules(iterable, необязательный) - итерируемый объект, содержащий пары (ключ, модуль), которыми будет инициализирован словарь. Если не указан, создаётся пустой словарь.
Пример инициализации и добавления модулей
Создадим пустой словарь модулей и добавим в него несколько линейных слоёв с разными ключами:
import torch
import torch.nn as nn
module_dict = nn.ModuleDict()
module_dict['layer1'] = nn.Linear(10, 20)
module_dict['layer2'] = nn.Linear(20, 30)
module_dict['layer3'] = nn.Linear(30, 10)
print(module_dict)
Результат выполнения кода:
ModuleDict(
(layer1): Linear(in_features=10, out_features=20, bias=True)
(layer2): Linear(in_features=20, out_features=30, bias=True)
(layer3): Linear(in_features=30, out_features=10, bias=True)
)
Пример инициализации с данными
Инициализируем ModuleDict сразу с несколькими модулями, передав список кортежей:
import torch
import torch.nn as nn
modules = [
('fc1', nn.Linear(5, 10)),
('fc2', nn.Linear(10, 5)),
]
module_dict = nn.ModuleDict(modules)
print(module_dict)
Результат выполнения кода:
ModuleDict(
(fc1): Linear(in_features=5, out_features=10, bias=True)
(fc2): Linear(in_features=10, out_features=5, bias=True)
)
Пример доступа к модулям и использования в forward
Покажем, как использовать модули из словаря в методе forward. Обратите внимание, что все модули автоматически регистрируются, и их параметры будут оптимизироваться:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.ModuleDict({
'linear1': nn.Linear(5, 10),
'linear2': nn.Linear(10, 1),
})
def forward(self, x):
x = torch.relu(self.layers['linear1'](x))
x = self.layers['linear2'](x)
return x
model = MyModel()
x = torch.randn(3, 5)
res = model(x)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 1])
Пример использования методов словаря
ModuleDict поддерживает основные методы стандартного словаря: keys, values, items, get, pop, update и другие. Рассмотрим их на примере:
import torch
import torch.nn as nn
module_dict = nn.ModuleDict()
module_dict['conv1'] = nn.Conv2d(3, 16, 3)
module_dict['conv2'] = nn.Conv2d(16, 32, 3)
# Получение ключей
print(list(module_dict.keys()))
# Получение значений
print(len(list(module_dict.values())))
# Проверка наличия ключа
print('conv1' in module_dict)
# Удаление модуля
module_dict.pop('conv1')
print(list(module_dict.keys()))
Результат выполнения кода:
['conv1', 'conv2']
2
True
['conv2']
Смотрите также
-
класс
ParameterList,
который хранит параметры в виде списка -
класс
ParameterDict,
который хранит параметры в виде словаря -
класс
UninitializedParameter,
который представляет неинициализированный параметр -
класс
Buffer,
который хранит буферы (необновляемые тензоры) в модуле