Spec-Zone.ru › PyTorch 2

torch.func.grad

torch.func.grad(func, argnums=0, has_aux=False)

grad оператор помогает вычислять градиенты func по отношению к входному(ым) аргументу(ам), указанным в argnums. Этот оператор может быть вложен для вычисления градиентов высших порядков.

Параметры
  • func (Callable) – Функция Python, принимающая один или несколько аргументов. Должна возвращать тензор с одним элементом. Если указано has_aux равно True, функция может возвращать кортеж из тензора с одним элементом и других вспомогательных объектов: (output, aux).
  • argnums (int или Кортеж[int]) – Указывает аргументы, по которым необходимо вычислить градиенты. argnums может быть целым числом или кортежем целых чисел. По умолчанию: 0.
  • has_aux (bool) – Флаг, указывающий, что func возвращает тензор и другие вспомогательные объекты: (output, aux). По умолчанию: False.
Возвращает

Функцию для вычисления градиентов по её входным данным. По умолчанию, выходной тензор функции является(ются) градиентом(ами) по первому аргументу. Если указано has_aux равно True, возвращается кортеж градиентов и вспомогательных объектов. Если argnums является кортежем целых чисел, возвращается кортеж выходных градиентов по каждому argnums значению.

Тип возвращаемого значения

Callable

Пример использования grad:

>>> from torch.func import grad
>>> x = torch.randn([])
>>> cos_x = grad(lambda x: torch.sin(x))(x)
>>> assert torch.allclose(cos_x, x.cos())
>>>
>>> # Second-order gradients
>>> neg_sin_x = grad(grad(lambda x: torch.sin(x)))(x)
>>> assert torch.allclose(neg_sin_x, -x.sin())

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

>>> from torch.func import grad, vmap
>>> batch_size, feature_size = 3, 5
>>>
>>> def model(weights, feature_vec):
>>>     # Very simple linear model with activation
>>>     assert feature_vec.dim() == 1
>>>     return feature_vec.dot(weights).relu()
>>>
>>> def compute_loss(weights, example, target):
>>>     y = model(weights, example)
>>>     return ((y - target) ** 2).mean()  # MSELoss
>>>
>>> weights = torch.randn(feature_size, requires_grad=True)
>>> examples = torch.randn(batch_size, feature_size)
>>> targets = torch.randn(batch_size)
>>> inputs = (weights, examples, targets)
>>> grad_weight_per_example = vmap(grad(compute_loss), in_dims=(None, 0, 0))(*inputs)

Пример использования grad с has_aux и argnums:

>>> from torch.func import grad
>>> def my_loss_func(y, y_pred):
>>>    loss_per_sample = (0.5 * y_pred - y) ** 2
>>>    loss = loss_per_sample.mean()
>>>    return loss, (y_pred, loss_per_sample)
>>>
>>> fn = grad(my_loss_func, argnums=(0, 1), has_aux=True)
>>> y_true = torch.rand(4)
>>> y_preds = torch.rand(4, requires_grad=True)
>>> out = fn(y_true, y_preds)
>>> # > output is ((grads w.r.t y_true, grads w.r.t y_preds), (y_pred, loss_per_sample))

Примечание

Использование PyTorch torch.no_grad вместе с grad.

Случай 1: Использование torch.no_grad внутри функции:

>>> def f(x):
>>>     with torch.no_grad():
>>>         c = x ** 2
>>>     return x - c

В этом случае, grad(f)(x) будет учитывать внутренний torch.no_grad.

Случай 2: Использование grad внутри контекстного менеджера torch.no_grad:

>>> with torch.no_grad():
>>>     grad(f)(x)

В этом случае, grad будет учитывать внутренний torch.no_grad, но не внешний. Это потому, что grad является «преобразованием функции»: её результат не должен зависеть от результата контекстного менеджера вне f.

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.grad.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API