Потомки модуля в PyTorch
Метод children перебирает
только прямых потомков модуля.
Вложенные слои внутри дочернего
блока в этот перебор не попадают.
На модели с двумя линейными
слоями на верхнем уровне
увидим ровно два элемента:
import torch
import torch.nn as nn
class TwoLayer(nn.Module):
def __init__(self):
super().__init__()
self.first = nn.Linear(2, 3)
self.second = nn.Linear(3, 1)
def forward(self, x):
return self.second(self.first(x))
net = TwoLayer()
kinds = [type(child).__name__ for child in net.children()]
print(kinds) # выведет ['Linear', 'Linear']
Такой обход удобен, когда нужно пройти только верхний уровень дерева, не заходя глубже:
import torch
import torch.nn as nn
class TwoLayer(nn.Module):
def __init__(self):
super().__init__()
self.first = nn.Linear(2, 3)
self.second = nn.Linear(3, 1)
def forward(self, x):
return self.second(self.first(x))
net = TwoLayer()
count = sum(1 for _ in net.children())
print(count) # выведет 2
Соберите модуль с линейными
слоями 2 на 4
и 4 на 1
как двух прямых потомков.
Выведите число элементов
в переборе непосредственных
детей корня.
Опишите блок с двумя линейными преобразованиями на верхнем уровне. Выведите список имён типов для каждого прямого потомка.
Создайте сеть из линейного
слоя 3 на 2
и второго 2 на 3
без общих контейнеров между
ними. Выведите, сколько раз
срабатывает перебор прямых
потомков у корневого модуля.