Метод modules
Метод modules класса Module возвращает итератор,
который обходит все дочерние модули модели, включая саму модель.
Этот метод полезен для рекурсивного обхода всех подслоев нейронной сети,
например, для инициализации параметров, применения преобразований
или анализа архитектуры модели.
Синтаксис
model.modules()
Пример
Давайте создадим простую модель и выведем все её модули:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(10, 20)
self.fc2 = nn.Linear(20, 5)
self.relu = nn.ReLU()
model = SimpleModel()
for module in model.modules():
print(module)
Результат выполнения кода:
SimpleModel(
(fc1): Linear(in_features=10, out_features=20, bias=True)
(fc2): Linear(in_features=20, out_features=5, bias=True)
(relu): ReLU()
)
Linear(in_features=10, out_features=20, bias=True)
Linear(in_features=20, out_features=5, bias=True)
ReLU()
Пример
Давайте применим метод modules для инициализации
весов всех линейных слоёв модели:
import torch
import torch.nn as nn
class ComplexModel(nn.Module):
def __init__(self):
super().__init__()
self.seq = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 30),
nn.ReLU(),
)
self.final = nn.Linear(30, 5)
model = ComplexModel()
def init_weights(module):
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
module.bias.data.fill_(0.0)
for module in model.modules():
init_weights(module)
print("Инициализация завершена")
Результат выполнения кода:
"Инициализация завершена"
Пример
Давайте соберём все параметры модели с использованием
modules и сравним с встроенным методом parameters:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(5, 3)
self.layer2 = nn.Linear(3, 2)
model = MyModel()
param_count = 0
for module in model.modules():
if isinstance(module, nn.Linear):
param_count += sum(p.numel() for p in module.parameters())
print(f"Всего параметров в линейных слоях: {param_count}")
total_params = sum(p.numel() for p in model.parameters())
print(f"Всего параметров в модели: {total_params}")
Результат выполнения кода:
"Всего параметров в линейных слоях: 28"
"Всего параметров в модели: 28"
Смотрите также
-
метод
named_modules,
который возвращает итератор с именами модулей -
метод
children,
который возвращает итератор по непосредственным дочерним модулям -
метод
apply,
который рекурсивно применяет функцию ко всем подмодулям -
метод
parameters,
который возвращает итератор по всем обучаемым параметрам