Метод apply
Метод apply класса Module рекурсивно применяет переданную функцию ко всем дочерним модулям и самому модулю.
Первый параметр метода - функция, которая принимает на вход модуль и выполняет над ним некоторое действие.
Метод возвращает исходный модуль, что удобно для цепочек вызовов.
Этот метод часто используется для инициализации весов, изменения типа параметров или применения общих преобразований ко всей модели.
Синтаксис
module.apply(fn)
Параметры
Метод apply принимает следующий параметр:
-
fn- функция, которая принимает один аргумент - дочерний модуль. Функция может изменять модуль на месте и не должна возвращать никакого значения. Если функция возвращает значение, оно игнорируется.
Пример
Давайте создадим простую линейную модель и применим функцию, которая выводит имя каждого модуля:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(10, 20)
self.layer2 = nn.Linear(20, 5)
model = MyModel()
def print_module_name(module):
print(module.__class__.__name__)
model.apply(print_module_name)
Результат выполнения кода:
MyModel
Linear
Linear
Пример
Метод apply часто используется для инициализации весов.
Давайте создадим модель и инициализируем все линейные слои с помощью функции:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(10, 20)
self.layer2 = nn.Linear(20, 5)
model = MyModel()
def init_weights(module):
if isinstance(module, nn.Linear):
torch.manual_seed(0)
nn.init.xavier_uniform_(module.weight)
nn.init.zeros_(module.bias)
model.apply(init_weights)
print(model.layer1.weight[0, :5])
Результат выполнения кода:
tensor([0.7526, 0.0790, 0.3313, 0.5112, 0.4344])
Пример
Метод apply можно использовать для изменения типа параметров всех модулей.
Давайте переведём все модули в тип float64:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(10, 20)
self.layer2 = nn.Linear(20, 5)
model = MyModel()
def to_double(module):
if hasattr(module, 'weight'):
module.weight.data = module.weight.data.double()
if hasattr(module, 'bias') and module.bias is not None:
module.bias.data = module.bias.data.double()
model.apply(to_double)
print(model.layer1.weight.dtype)
print(model.layer2.weight.dtype)
Результат выполнения кода:
torch.float64
torch.float64
Пример
Метод apply можно использовать для заморозки параметров градиента в выбранных модулях.
Давайте заморозим все слои, кроме последнего:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(10, 20)
self.layer2 = nn.Linear(20, 5)
self.final = nn.Linear(5, 1)
model = MyModel()
def freeze_weights(module):
if hasattr(module, 'weight') and module is not model.final:
module.weight.requires_grad = False
if hasattr(module, 'bias') and module.bias is not None and module is not model.final:
module.bias.requires_grad = False
model.apply(freeze_weights)
print(model.layer1.weight.requires_grad)
print(model.final.weight.requires_grad)
Результат выполнения кода:
False
True
Смотрите также
-
класс
Module,
который является базовым классом для всех модулей в PyTorch -
метод
children,
который возвращает итератор по дочерним модулям -
метод
modules,
который рекурсивно обходит все модули в модели -
метод
parameters,
который возвращает итератор по параметрам модуля