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

Атрибут is_leaf

Атрибут is_leaf класса Tensor представляет собой булево значение, которое показывает, является ли тензор листовым (leaf) в графе вычислений. Листовые тензоры - это тензоры, которые не были созданы в результате каких-либо операций, то есть они были созданы пользователем явно или получены как параметры модуля. Этот атрибут играет важную роль при автоматическом дифференцировании, так как градиенты вычисляются и сохраняются только для листовых тензоров, если для них установлен флаг requires_grad.

Синтаксис

tensor.is_leaf

Пример

Давайте создадим тензор и проверим его атрибут is_leaf:

import torch t = torch.tensor([1, 2, 3, 4, 5]) print(t.is_leaf)

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

True

Пример

Теперь создадим тензор с requires_grad и выполним над ним операцию, затем проверим is_leaf для нового тензора:

import torch t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], requires_grad=True) t2 = t * 2 print(t.is_leaf) print(t2.is_leaf)

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

True False

Как видно из примера, тензор t является листовым, так как он создан явно, а тензор t2, полученный в результате операции, листовым не является.

Пример

Рассмотрим ситуацию с параметрами модели. Параметры слоя Linear являются листовыми тензорами:

import torch import torch.nn as nn layer = nn.Linear(5, 3) print(layer.weight.is_leaf) print(layer.bias.is_leaf)

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

True True

Пример

Атрибут is_leaf также можно использовать для проверки, сохраняется ли градиент для тензора. Градиент сохраняется только для листовых тензоров, у которых requires_grad равен True:

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

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

tensor([2., 2., 2., 2., 2.]) None

В этом примере градиент был вычислен для листового тензора t, а для нелистового t2 градиент не сохранился (равен None).

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

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