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

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