Метод named_modules
Метод named_modules класса Module
возвращает генератор, который перебирает все
подмодули текущего модуля, включая его самого.
Каждый элемент генератора представляет собой
кортеж из двух элементов: имени модуля (строки)
и самого модуля (объекта Module).
Синтаксис
module.named_modules([memo=None, prefix=''])
Необязательный параметр memo используется
для внутренней работы метода и обычно не
передается пользователем. Параметр prefix
задает префикс для имен модулей.
Пример
Давайте создадим простую модель и выведем
имена всех модулей с помощью метода named_modules:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(3, 16, 3)
self.relu = nn.ReLU()
self.pool = nn.MaxPool2d(2)
model = SimpleModel()
for name, module in model.named_modules():
print(f"{name}: {module.__class__.__name__}")
Результат выполнения кода:
"": SimpleModel
conv: Conv2d
relu: ReLU
pool: MaxPool2d
Пример
Метод named_modules позволяет получить
доступ к каждому модулю и его атрибутам, например,
к весам:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(10, 5)
self.activation = nn.ReLU()
model = SimpleModel()
for name, module in model.named_modules():
if hasattr(module, 'weight'):
print(f"{name}: {module.weight.shape}")
Результат выполнения кода:
"": torch.Size([5, 10])
linear: torch.Size([5, 10])
Пример
С помощью метода named_modules можно
изменять атрибуты всех модулей, например,
перевести их в режим обучения или оценки:
import torch
import torch.nn as nn
class ComplexModel(nn.Module):
def __init__(self):
super().__init__()
self.block = nn.Sequential(
nn.Linear(20, 10),
nn.ReLU(),
nn.Linear(10, 5)
)
model = ComplexModel()
model.eval()
for name, module in model.named_modules():
print(f"{name}: training={module.training}")
Результат выполнения кода:
"": training=False
block: training=False
block.0: training=False
block.1: training=False
block.2: training=False
Пример
Метод named_modules можно использовать
для сбора информации о всех модулях модели
в словарь:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(128, 64)
self.fc2 = nn.Linear(64, 10)
self.dropout = nn.Dropout(0.5)
model = MyModel()
module_dict = {}
for name, module in model.named_modules():
module_dict[name] = module.__class__.__name__
print(module_dict)
Результат выполнения кода:
{'': 'MyModel', 'fc1': 'Linear', 'fc2': 'Linear', 'dropout': 'Dropout'}
Смотрите также
-
метод
modules,
который возвращает генератор всех модулей без имен -
метод
named_children,
который возвращает только прямые дочерние модули -
метод
children,
который возвращает итератор прямых дочерних модулей -
метод
named_parameters,
который возвращает имена и параметры модели