PyTorch torch.no_grad Function
PyTorch torch Reference Manual
torch.no_gradis a context manager in PyTorch used to disable gradient computation. Inno_gradTensors created inside the block do not compute gradients, which can significantly reduce memory consumption and improve inference speed.
This is essential during model inference and evaluation stages, greatly improving performance and saving memory.
Function Definition
torch.no_grad()
Parameters:
- No parameters. This is a context manager.
Return Value:
- Returns a context manager that disables gradient computation in that context.
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)
# Compute inside the no_grad context
with torch.no_grad():
y = x * 2
print("Inside no_grad:", y.requires_grad)
# Outside no_grad
z = x * 2
print("Outside no_grad:", z.requires_grad)
# Create a tensor that requires gradients
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
# Compute inside the no_grad context
with torch.no_grad():
y = x * 2
print("Inside no_grad:", y.requires_grad)
# Outside no_grad
z = x * 2
print("Outside no_grad:", z.requires_grad)
The output is:
在 no_grad 中: False 在 no_grad 外: True
Example 2: Model Inference
Example
import torch
import torch.nn as nn
# Define a simple model
model = nn.Linear(10, 2)
# Create input data
x = torch.randn(1, 10)
# Use no_grad during inference
with torch.no_grad():
output = model(x)
print("Output:", output)
# No gradient computation required
# Or use the decorator
# @torch.no_grad()
# def predict(x):
# return model(x)
import torch.nn as nn
# Define a simple model
model = nn.Linear(10, 2)
# Create input data
x = torch.randn(1, 10)
# Use no_grad during inference
with torch.no_grad():
output = model(x)
print("Output:", output)
# No gradient computation required
# Or use the decorator
# @torch.no_grad()
# def predict(x):
# return model(x)
The output is:
输出: tensor([[0.0920, 0.3557]])
Example 3: Evaluation Mode
Example
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(10, 5),
nn.ReLU(),
nn.Linear(5, 2)
)
# Switch to evaluation mode
model.eval()
# Prepare test data
test_input = torch.randn(5, 10)
# Disable gradients during evaluation
with torch.no_grad():
predictions = model(test_input)
print("Prediction result shape:", predictions.shape)
import torch.nn as nn
model = nn.Sequential(
nn.Linear(10, 5),
nn.ReLU(),
nn.Linear(5, 2)
)
# Switch to evaluation mode
model.eval()
# Prepare test data
test_input = torch.randn(5, 10)
# Disable gradients during evaluation
with torch.no_grad():
predictions = model(test_input)
print("Prediction result shape:", predictions.shape)
The output is:
预测结果形状: torch.Size([5, 2])
Example 4: Comparing Memory Usage
Example
import torch
import torch.nn as nn
model = nn.Linear(1000, 1000)
# Create a large amount of input
inputs = [torch.randn(100, 1000) for _ in range(100)]
# Without no_grad (gradient history will be recorded)
print("Without no_grad:")
for inp in inputs[:5]:
_ = model(inp)
# Use no_grad (does not record gradient history)
print("With no_grad:")
with torch.no_grad():
for inp in inputs[:5]:
_ = model(inp)
import torch.nn as nn
model = nn.Linear(1000, 1000)
# Create a large amount of input
inputs = [torch.randn(100, 1000) for _ in range(100)]
# Without no_grad (gradient history will be recorded)
print("Without no_grad:")
for inp in inputs[:5]:
_ = model(inp)
# Use no_grad (does not record gradient history)
print("With no_grad:")
with torch.no_grad():
for inp in inputs[:5]:
_ = model(inp)
Usingno_gradcan significantly reduce memory consumption because there is no need to store gradient information for intermediate variables.
Related Functions
torch.enable_grad(): Enables gradient computation (as opposed tono_gradthe opposite).torch.set_grad_enabled(grad): Enables or disables gradient computation based on parameters.torch.inference_mode(): A stricter inference mode that disables both gradient computation and autograd.
Notes
- During model inference and evaluation, always use
torch.no_grad()to save memory and improve speed. - If
no_gradvariables created inside the block need to be used outside the block, they need to be manually copied out. - and
model.eval()Use in combination for best results.
Other Extensions