enable_grad
-
class torch.enable_grad[source] -
Менеджер контекста, который включает вычисление градиента.
Включает вычисление градиента, если оно было отключено с помощью
no_gradилиset_grad_enabled.Этот менеджер контекста локален для потока; он не повлияет на вычисления в других потоках.
Также работает как декоратор. (Убедитесь, что вы инициализируете его с круглыми скобками.)
Примечание
enable_grad — один из нескольких механизмов, которые могут локально включить или отключить градиенты. Смотрите Локальное отключение вычисления градиента для получения дополнительной информации о том, как они сравниваются.
Примечание
Этот API не применим к 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
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.enable_grad.html