Класс ModuleList
Класс ModuleList предназначен для хранения списка модулей в PyTorch.
Он позволяет объединять несколько подслоев в один контейнер и управлять ими как единым целым.
В отличие от обычного списка Python, ModuleList автоматически регистрирует все добавленные модули,
что позволяет параметрам этих модулей корректно отслеживаться при обучении модели.
Синтаксис
torch.nn.ModuleList(modules=None)
Конструктор класса принимает один необязательный параметр:
- ⁅i⁆modules⁅/i⁆ (итерируемый объект, необязательно) - последовательность модулей, которые будут добавлены в список.
Пример
Создадим список из двух линейных слоёв с помощью ModuleList:
import torch
import torch.nn as nn
# Создаём ModuleList с двумя линейными слоями
layers = nn.ModuleList([
nn.Linear(10, 20),
nn.Linear(20, 30)
])
print(layers)
Результат выполнения кода:
ModuleList(
(0): Linear(in_features=10, out_features=20, bias=True)
(1): Linear(in_features=20, out_features=30, bias=True)
)
Пример
Рассмотрим создание пустого списка и его последующее заполнение с помощью метода append:
import torch
import torch.nn as nn
# Создаём пустой ModuleList
layers = nn.ModuleList()
# Добавляем модули с помощью метода append
layers.append(nn.Linear(5, 10))
layers.append(nn.ReLU())
layers.append(nn.Linear(10, 1))
print(layers)
Результат выполнения кода:
ModuleList(
(0): Linear(in_features=5, out_features=10, bias=True)
(1): ReLU()
(2): Linear(in_features=10, out_features=1, bias=True)
)
Пример
Продемонстрируем работу с ModuleList внутри пользовательского модуля.
Список слоёв используется как часть архитектуры нейронной сети:
import torch
import torch.nn as nn
class MyNet(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.ModuleList([
nn.Linear(3, 4),
nn.ReLU(),
nn.Linear(4, 2)
])
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x
# Создаём экземпляр сети
model = MyNet()
# Пример входных данных
t = torch.randn(1, 3)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 2])
Смотрите также
-
класс
ModuleList,
который предоставляет контейнер для хранения последовательности модулей -
метод
append,
который добавляет модуль в конец списка -
метод
extend,
который добавляет несколько модулей в конец списка -
метод
insert,
который вставляет модуль на указанную позицию