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

Метод requires_grad_

Метод requires_grad_ класса Module изменяет флаг requires_grad для всех параметров модели. Первым параметром метод принимает булево значение requires_grad, которое определяет, нужно ли вычислять градиенты для параметров модуля. Если передать True, градиенты будут вычисляться, если False - вычисление градиентов будет отключено.

Синтаксис

module.requires_grad_(requires_grad=True)

Пример

Давайте создадим простую линейную модель и проверим состояние параметров по умолчанию:

import torch model = torch.nn.Linear(3, 1) for param in model.parameters(): print(param.requires_grad)

Результат выполнения кода:

True True

Как видим, по умолчанию все параметры требуют вычисления градиентов.

Пример

Теперь отключим вычисление градиентов для всех параметров модели:

import torch model = torch.nn.Linear(3, 1) model.requires_grad_(False) for param in model.parameters(): print(param.requires_grad)

Результат выполнения кода:

False False

Пример

Включим вычисление градиентов обратно для всех параметров:

import torch model = torch.nn.Linear(3, 1) model.requires_grad_(False) model.requires_grad_(True) for param in model.parameters(): print(param.requires_grad)

Результат выполнения кода:

True True

Пример

Рассмотрим практическое применение метода при заморозке слоев модели:

import torch class MyModel(torch.nn.Module): def __init__(self): super().__init__() self.fc1 = torch.nn.Linear(10, 5) self.fc2 = torch.nn.Linear(5, 2) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) model = MyModel() model.fc1.requires_grad_(False) print("fc1 requires_grad:", model.fc1.weight.requires_grad) print("fc2 requires_grad:", model.fc2.weight.requires_grad)

Результат выполнения кода:

fc1 requires_grad: False fc2 requires_grad: True

Это полезно при тонкой настройке (fine-tuning) предобученных моделей, когда нужно заморозить некоторые слои.

Смотрите также

  • метод parameters,
    который возвращает итератор по параметрам модуля
  • метод train,
    который переводит модуль в режим обучения
  • метод eval,
    который переводит модуль в режим оценки
  • атрибут training,
    который указывает, находится ли модуль в режиме обучения
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить