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

Метод requires_grad_

Метод requires_grad_ класса Tensor изменяет значение атрибута requires_grad тензора. Этот атрибут определяет, будет ли тензор отслеживать все операции, выполняемые над ним, для последующего вычисления градиентов с помощью метода backward. Метод возвращает сам тензор, что позволяет использовать его в цепочках вызовов.

Синтаксис

tensor.requires_grad_(requires_grad=True)

Единственным параметром метода является булево значение requires_grad, которое указывает, нужно ли включать отслеживание градиентов. По умолчанию этот параметр равен True.

Пример

Создадим обычный тензор без отслеживания градиентов и включим эту возможность с помощью метода requires_grad_:

import torch t = torch.tensor([1.0, 2.0, 3.0]) print(t.requires_grad) t.requires_grad_(True) print(t.requires_grad)

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

False True

Как видно, после применения метода атрибут requires_grad изменился на True.

Пример

Метод часто используется для отключения отслеживания градиентов у тензоров, которые не должны участвовать в процессе обучения, например, для заморозки параметров модели. В этом примере мы выключим отслеживание градиентов:

import torch t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) print(t.requires_grad) t.requires_grad_(False) print(t.requires_grad)

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

True False

Теперь тензор не отслеживает операции, и вызов backward на нем или на результатах вычислений с его участием будет невозможен.

Пример

Рассмотрим практический пример, где метод requires_grad_ используется для включения отслеживания градиентов у тензора весов перед вычислением градиентов:

import torch # Создаем тензор весов без градиента weights = torch.ones(3) print("До:", weights.requires_grad) # Включаем отслеживание градиентов weights.requires_grad_(True) print("После:", weights.requires_grad) # Выполняем операцию x = torch.tensor([1.0, 2.0, 3.0]) loss = (weights * x).sum() # Вычисляем градиенты loss.backward() print(weights.grad)

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

До: False После: True tensor([1., 2., 3.])

После включения отслеживания градиентов мы смогли вычислить градиент weights.grad.

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

  • атрибут requires_grad,
    который показывает, отслеживает ли тензор операции для градиентов
  • атрибут grad,
    который содержит вычисленный градиент тензора
  • атрибут grad_fn,
    который указывает функцию, создавшую тензор
  • метод backward,
    который вычисляет градиенты для тензора
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить