Метод extend
Метод extend класса ModuleList предназначен для расширения списка модулей
путем добавления нескольких подмодулей из итерируемого объекта. В качестве
параметра метод принимает итерируемый объект (список, кортеж и т.д.),
содержащий модули nn.Module. Все добавляемые модули становятся
подмодулями текущего списка и регистрируются в нем для корректной
работы механизма оптимизации и сохранения состояния.
Синтаксис
module_list.extend(modules)
Где:
-
module_list- объект классаModuleList -
modules- итерируемый объект с модулями для добавления
Пример
Давайте создадим список модулей и расширим его несколькими линейными слоями:
import torch
import torch.nn as nn
# Создаем ModuleList с одним слоем
layers = nn.ModuleList([nn.Linear(10, 20)])
print(f"Количество слоев: {len(layers)}")
# Расширяем список тремя слоями
layers.extend([
nn.ReLU(),
nn.Linear(20, 30),
nn.ReLU()
])
print(f"Количество слоев: {len(layers)}")
Результат выполнения кода:
Количество слоев: 1
Количество слоев: 4
Пример
Расширим список модулями, созданными с помощью генератора списка:
import torch
import torch.nn as nn
# Создаем пустой ModuleList
model_layers = nn.ModuleList()
# Расширяем с помощью генератора
model_layers.extend([
nn.Linear(5, 10) for _ in range(3)
])
# Выводим информацию о каждом слое
for i, layer in enumerate(model_layers):
print(f"Слой {i}: {layer}")
Результат выполнения кода:
Слой 0: Linear(in_features=5, out_features=10, bias=True)
Слой 1: Linear(in_features=5, out_features=10, bias=True)
Слой 2: Linear(in_features=5, out_features=10, bias=True)
Пример
Расширим список модулями из другого списка и проверим все подмодули:
import torch
import torch.nn as nn
# Создаем два ModuleList
list1 = nn.ModuleList([nn.Linear(10, 20)])
list2 = nn.ModuleList([nn.ReLU(), nn.Linear(20, 10)])
# Расширяем первый список вторым
list1.extend(list2)
# Выводим все модули
for i, module in enumerate(list1):
print(f"Модуль {i}: {module}")
Результат выполнения кода:
Модуль 0: Linear(in_features=10, out_features=20, bias=True)
Модуль 1: ReLU()
Модуль 2: Linear(in_features=20, out_features=10, bias=True)
Смотрите также
-
класс
ModuleList,
который представляет список модулей -
метод
append,
который добавляет один модуль в конец списка -
метод
insert,
который вставляет модуль по указанному индексу -
метод
extend,
который расширяет список несколькими модулями