set_grad_enabled
-
class torch.set_grad_enabled(mode)[source] -
Менеджер контекста, который включает или выключает вычисление градиента.
set_grad_enabledбудет включать или отключать градиенты, основываясь на его аргументеmode. Он может использоваться как менеджер контекста или как функция.Этот менеджер контекста является локальным для потока; он не повлияет на вычисления в других потоках.
- Параметры
-
mode (bool) – Флаг для включения вычисления градиентов (
True) или отключения (False). Это может использоваться для условного включения градиентов.
Примечание
set_grad_enabled — один из нескольких механизмов, которые могут включать или выключать градиенты локально. См. Локальное отключение вычисления градиентов для получения дополнительной информации о том, как они сравниваются.
Примечание
Этот API не применяется к forward-mode 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/2.1/generated/torch.set_grad_enabled.html