Метод detach
Метод detach класса Tensor создает новый тензор,
который разделяет данные с исходным тензором, но не имеет
истории вычислений и не отслеживает градиенты. Это позволяет
использовать данные тензора для дальнейших операций без участия
в автоматическом дифференцировании. Метод не принимает параметров
и возвращает новый тензор с теми же данными, но с отключенным
требованием градиента.
Синтаксис
tensor.detach()
Пример
Давайте создадим тензор с требованием градиента и отсоединим его:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
detached_t = t.detach()
print(t)
print(detached_t)
print(t.requires_grad)
print(detached_t.requires_grad)
Результат выполнения кода:
tensor([1., 2., 3.], requires_grad=True)
tensor([1., 2., 3.])
True
False
Пример
Покажем, что отсоединенный тензор не участвует в вычислении градиентов:
import torch
torch.manual_seed(0)
t = torch.tensor([1., 2., 3.], requires_grad=True)
detached_t = t.detach()
res = detached_t.sum() + t.sum()
res.backward()
print(t.grad)
print(detached_t.grad)
Результат выполнения кода:
tensor([1., 1., 1.])
None
Как видно из примера, градиент вычисляется только для исходного тензора t,
а для отсоединенного тензора detached_t градиент отсутствует.
Пример
Используем метод detach для получения данных тензора
в виде массива NumPy без сохранения истории вычислений:
import torch
import numpy as np
torch.manual_seed(0)
t = torch.randn(3, requires_grad=True)
numpy_array = t.detach().numpy()
print(numpy_array)
print(type(numpy_array))
Результат выполнения кода:
[ 1.541 -0.2934 -2.1788]
<class 'numpy.ndarray'>
Пример
Метод detach часто используется при обучении моделей
для отсоединения промежуточных результатов. Например, при
использовании метода detach в цикле обучения:
import torch
torch.manual_seed(0)
t = torch.randn(3, requires_grad=True)
loss = (t ** 2).sum()
# Отсоединяем тензор для использования в других операциях
detached_loss = loss.detach()
print(loss)
print(detached_loss)
print(loss.requires_grad)
print(detached_loss.requires_grad)
Результат выполнения кода:
tensor(7.3105, grad_fn=<SumBackward0>)
tensor(7.3105)
True
False