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

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