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

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