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

Атрибут grad

Атрибут grad класса Tensor содержит градиент тензора, накопленный в ходе обратного распространения ошибки. Значение атрибута доступно только для тензоров, у которых флаг requires_grad установлен в True. Если градиент ещё не был вычислен, атрибут возвращает None. Градиент представляет собой тензор той же формы, что и исходный тензор, и содержит частные производные по каждому элементу.

Синтаксис

tensor.grad

Пример

Давайте вычислим градиент для простой функции:

import torch t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) res = (t ** 2).sum() res.backward() print(t.grad)

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

tensor([2., 4., 6.])

Градиент содержит производные функции sum(t^2) по каждому элементу: 2*t.

Пример

Атрибут grad можно использовать для доступа к градиенту в процессе обучения модели:

import torch w = torch.tensor([0.5, -0.2, 0.8], requires_grad=True) x = torch.tensor([1.0, 2.0, 3.0]) y_true = torch.tensor([4.0]) y_pred = (w * x).sum() loss = (y_pred - y_true) ** 2 loss.backward() print(w.grad)

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

tensor([-0.8000, -1.6000, -2.4000])

Пример

Градиент накапливается при повторных вызовах backward. Чтобы избежать накопления, нужно обнулять градиент вручную:

import torch torch.manual_seed(0) w = torch.randn(3, requires_grad=True) x = torch.tensor([1.0, 2.0, 3.0]) for _ in range(2): loss = (w * x).sum() loss.backward() print(w.grad) w.grad.zero_()

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

tensor([1., 2., 3.]) tensor([1., 2., 3.])

Пример

Проверим, вычислен ли градиент для тензора:

import torch t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) print(t.grad is None) (t ** 2).sum().backward() print(t.grad is None)

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

True False

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

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