Spec-Zone.ru › PyTorch 1

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

Spec-Zone.ru

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