Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API