torch.func.grad
-
torch.func.grad(func, argnums=0, has_aux=False)[source] -
Оператор
gradпомогает вычислять градиентыfuncотносительно входных данных, указанных вargnums. Этот оператор можно вкладывать друг в друга для вычисления градиентов более высоких порядков.- Параметры:
-
-
func (Callable) – Функция Python, принимающая один или несколько аргументов. Должна возвращать тензор с одним элементом. Если указано, что
has_auxравноTrue, функция может возвращать кортеж из тензора с одним элементом и других вспомогательных объектов:(output, aux). -
argnums (int или Tuple[int]) – Указывает аргументы, относительно которых вычисляются градиенты.
argnumsможет быть целым числом или кортежем целых чисел. По умолчанию: 0. -
has_aux (bool) – Флаг, указывающий, что
funcвозвращает тензор и другие вспомогательные объекты:(output, aux). По умолчанию: False.
-
func (Callable) – Функция Python, принимающая один или несколько аргументов. Должна возвращать тензор с одним элементом. Если указано, что
- Возвращает:
-
Функцию для вычисления градиентов относительно её входных данных. По умолчанию функция возвращает тензор(ы) градиента относительно первого аргумента. Если указано, что
has_auxравноTrue, возвращается кортеж из градиентов и вспомогательных выходных объектов. Еслиargnumsявляется кортежем целых чисел, возвращается кортеж выходных градиентов относительно каждого значенияargnums. - Тип возвращаемого значения:
Пример использования
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))
Примечание
Использование
torch.no_gradPyTorch вместе с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.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.func.grad.html