Функция detach
Функция detach создает новый тензор, который разделяет данные
с исходным тензором, но при этом не участвует в графе вычислений.
Это означает, что операции с отсоединенным тензором не будут
отслеживаться для автоматического дифференцирования.
Функция не принимает никаких параметров и вызывается как метод тензора.
Синтаксис
tensor.detach()
Пример
Давайте создадим тензор с включенным вычислением градиента и отсоединим его от графа:
import torch
t = torch.tensor([1, 2, 3, 4, 5], dtype=torch.float, requires_grad=True)
res = t.detach()
print(res)
print(res.requires_grad)
Результат выполнения кода:
tensor([1., 2., 3., 4., 5.])
False
Пример
Важно понимать, что отсоединенный тензор разделяет данные с исходным тензором. Изменения в одном тензоре отразятся на другом:
import torch
t = torch.tensor([1, 2, 3, 4, 5], dtype=torch.float, requires_grad=True)
detached = t.detach()
detached[0] = 100
print(t)
print(detached)
Результат выполнения кода:
tensor([100., 2., 3., 4., 5.], requires_grad=True)
tensor([100., 2., 3., 4., 5.])
Пример
Отсоединенный тензор не участвует в распространении градиента. Рассмотрим это на примере:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
detached = t.detach()
res = detached.sum()
res.backward()
print(t.grad)
Результат выполнения кода:
None
Пример
Часто detach используется для получения данных тензора
в виде массива NumPy или для выполнения операций, которые не должны
влиять на градиент:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
detached = t.detach()
# Операции с detached не отслеживаются
res = detached * 2
print(res)
print(res.requires_grad)
Результат выполнения кода:
tensor([2., 4., 6.])
False