enable_grad
-
class torch.enable_grad(orig_func=None)[source] -
Менеджер контекста, который включает вычисление градиента.
Включает вычисление градиента, если оно было отключено с помощью
no_gradилиset_grad_enabled.Этот менеджер контекста является локальным для потока; он не повлияет на вычисления в других потоках.
Также может использоваться как декоратор.
Примечание
enable_grad — один из нескольких механизмов, которые могут включать или отключать градиенты локально. См. Локальное отключение вычисления градиента для получения дополнительной информации о том, как они сопоставляются.
Примечание
Этот API не применяется к forward-mode AD.
- Пример::
-
>>> x = torch.tensor([1.], requires_grad=True) >>> with torch.no_grad(): ... with torch.enable_grad(): ... y = x * 2 >>> y.requires_grad True >>> y.backward() >>> x.grad tensor([2.]) >>> @torch.enable_grad() ... def doubler(x): ... return x * 2 >>> with torch.no_grad(): ... z = doubler(x) >>> z.requires_grad True >>> @torch.enable_grad ... def tripler(x): ... return x * 3 >>> with torch.no_grad(): ... z = tripler(x) >>> z.requires_grad True
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.enable_grad.html