PyTorch torch.set_grad_enabled Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.set_grad_enabledis a function in PyTorch used to globally set whether gradient computation is enabled. It can dynamically enable or disable gradient computation; unlike context managers, it is a function that can change the global state.

This is very useful when you need to dynamically control gradient computation based on conditions.

Function Definition

torch.set_grad_enabled(mode)

Parameters:

  • mode(bool): If it isTrue, enable gradient computation; if it isFalse, disable gradient computation.

Return Value:

  • Returns a context manager that can be used inwithstatements.

Usage Examples

Example 1: Basic Usage

Example

import torch

# Create a tensor that requires gradients
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)

# Disable gradient computation
torch.set_grad_enabled(False)
y1 = x * 2
print("After disabling gradients:", y1.requires_grad)

# Enable gradient computation
torch.set_grad_enabled(True)
y2 = x * 2
print("After enabling gradients:", y2.requires_grad)

The output is:

禁用梯度后: False
启用梯度后: True

Example 2: Using as Context Manager

Example

import torch

x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)

# Use a context manager
with torch.set_grad_enabled(False):
    y1 = x * 2
    print("Inside the context manager:", y1.requires_grad)

y2 = x * 2
print("Outside the context manager:", y2.requires_grad)

The output is:

在上下文管理器内: False
在上下文管理器外: True

Example 3: Dynamically Controlling Training and Inference

Example

import torch
import torch.nn as nn

model = nn.Linear(10, 2)

def forward_pass(x, training=True):
    """Control gradients based on the training parameter"""
    with torch.set_grad_enabled(training):
        output = model(x)
        print(f"training={training}, requires_grad={output.requires_grad}")
    return output

# Training mode
x = torch.randn(5, 10)
forward_pass(x, training=True)

# Inference mode
forward_pass(x, training=False)

The output is:

training=True, requires_grad=True
training=False, requires_grad=False

Example 4: Saving and Restoring Gradient State

Example

import torch

# Initial state
print("Initial state:", torch.is_grad_enabled())

# Create a context manager that returns to the original state
old = torch.is_grad_enabled()

# Temporarily disable gradients
with torch.set_grad_enabled(False):
    print("Inside:", torch.is_grad_enabled())

# Automatically restore (but here we need to manually restore)
torch.set_grad_enabled(old)
print("After restoration:", torch.is_grad_enabled())

The output is:

初始状态: True
在内部: False
恢复后: True

Related Functions

  • torch.no_grad(): Context manager that disables gradient computation.
  • torch.enable_grad(): Context manager that enables gradient computation.
  • torch.is_grad_enabled(): Checks whether gradient computation is currently enabled.

Notes

  • set_grad_enabledIt can be called directly as a function to change the global state, or used as a context manager.
  • When called as a function, you need to manually restore the original state, otherwise it will affect subsequent code.
  • It is recommended to use the context manager approach to ensure the state is correctly restored.
  • Be aware of the impact on the global state and use it carefully in complex code.

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions