set_grad_enabled
-
class torch.set_grad_enabled(mode)[source] -
Менеджер контекста, который включает или отключает вычисление градиентов.
set_grad_enabledбудет включать или отключать градиенты в зависимости от его аргументаmode. Он может использоваться как менеджер контекста или как функция.Этот менеджер контекста локален для потока; он не повлияет на вычисления в других потоках.
- Параметры:
-
mode (bool) – Флаг, указывающий, нужно ли включить (
True) или отключить (False). Это можно использовать для условного включения градиентов.
Примечание
set_grad_enabled — это один из нескольких механизмов, которые могут включить или отключить градиенты локально. См. Локальное отключение вычисления градиента для получения дополнительной информации о том, как они сравниваются.
Примечание
Этот API не применим к AD в режиме прямого прохода.
- Пример::
-
>>> x = torch.tensor([1.], requires_grad=True) >>> is_train = False >>> with torch.set_grad_enabled(is_train): ... y = x * 2 >>> y.requires_grad False >>> _ = torch.set_grad_enabled(True) >>> y = x * 2 >>> y.requires_grad True >>> _ = torch.set_grad_enabled(False) >>> y = x * 2 >>> y.requires_grad False
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.set_grad_enabled.html