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

Метод 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,
    который расширяет список несколькими модулями
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить