PyTorch torch.enable_grad Function
Pytorch torch Reference Manual
torch.enable_gradIt is a context manager in PyTorch used to enable gradient computation. It istorch.no_gradthe opposite, and is used to explicitly enable gradient computation in code blocks that require gradients.
This is very useful when you temporarily need to train a model in inference code, or when you need to enable gradient computation in a local region.
Function Definition
torch.enable_grad()
Parameters:
- No parameters. This is a context manager.
Return Value:
- Returns a context manager that enables gradient computation in that context.
Usage Examples
Example 1: Basic Usage
Example
import torch
# By default, in a no_grad context
with torch.no_grad():
x = torch.tensor([1.0, 2.0, 3.0])
print("In no_grad:", x.requires_grad)
# Use enable_grad to temporarily enable gradients
with torch.enable_grad():
y = x * 2
print("In enable_grad:", y.requires_grad)
# After exiting, restore no_grad state
z = x * 2
print("After exiting:", z.requires_grad)
# By default, in a no_grad context
with torch.no_grad():
x = torch.tensor([1.0, 2.0, 3.0])
print("In no_grad:", x.requires_grad)
# Use enable_grad to temporarily enable gradients
with torch.enable_grad():
y = x * 2
print("In enable_grad:", y.requires_grad)
# After exiting, restore no_grad state
z = x * 2
print("After exiting:", z.requires_grad)
The output is:
在 no_grad 中: False 在 enable_grad 中: True 退出后: False
Example 2: Mixing Training and Inference
Example
import torch
import torch.nn as nn
model = nn.Linear(10, 2)
# Inference mode
with torch.no_grad():
x = torch.randn(5, 10)
output1 = model(x)
print("Inference output:", output1.shape)
# If you need to temporarily train some parameters during inference
with torch.enable_grad():
# Create a tensor that requires gradients for computation
temp_weight = torch.randn(10, 10, requires_grad=True)
temp_output = torch.mm(x, temp_weight)
print("Temporarily enable gradients:", temp_output.requires_grad)
import torch.nn as nn
model = nn.Linear(10, 2)
# Inference mode
with torch.no_grad():
x = torch.randn(5, 10)
output1 = model(x)
print("Inference output:", output1.shape)
# If you need to temporarily train some parameters during inference
with torch.enable_grad():
# Create a tensor that requires gradients for computation
temp_weight = torch.randn(10, 10, requires_grad=True)
temp_output = torch.mm(x, temp_weight)
print("Temporarily enable gradients:", temp_output.requires_grad)
The output is:
推理输出: torch.Size([5, 2]) 临时启用梯度: True
Example 3: Using as a Decorator
Example
import torch
@torch.enable_grad()
def train_step(x, y):
"""Simulate training steps"""
# This function will enable gradient computation internally
loss = (x - y).sum()
return loss
# Call it in a no_grad context
with torch.no_grad():
x = torch.tensor([1.0, 2.0, 3.0])
y = torch.tensor([0.0, 0.0, 0.0])
# The decorator ensures gradient computation is enabled
loss = train_step(x, y)
print("Loss:", loss)
print("Requires gradient:", loss.requires_grad)
@torch.enable_grad()
def train_step(x, y):
"""Simulate training steps"""
# This function will enable gradient computation internally
loss = (x - y).sum()
return loss
# Call it in a no_grad context
with torch.no_grad():
x = torch.tensor([1.0, 2.0, 3.0])
y = torch.tensor([0.0, 0.0, 0.0])
# The decorator ensures gradient computation is enabled
loss = train_step(x, y)
print("Loss:", loss)
print("Requires gradient:", loss.requires_grad)
The output is:
Loss: tensor(6.) 需要梯度: True
Related Functions
torch.no_grad(): Disables gradient computation.torch.set_grad_enabled(grad): Enables or disables gradient computation according to parameters.torch.is_grad_enabled(): Checks whether gradient computation is currently enabled.
Notes
enable_gradMainly used tono_gradtemporarily enable gradient computation in a context.- If used in an environment where gradients are globally enabled
enable_grad, it will have no effect. - It is recommended to use
enable_gradthe decorator to ensure that gradients are always enabled inside the function.
Other Extensions