Метод 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,
который указывает, находится ли модуль в режиме обучения