PyTorch torch.set_grad_enabled Function
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 in
withstatements.
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)
# 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)
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)
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())
# 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.
Other Extensions