PyTorch torch.inference_mode Function
Pytorch torch Reference Manual
torch.inference_modeis a context manager in PyTorch used for inference mode. It is moretorch.no_gradstrict, not only disabling gradient computation, but also disabling all tracking functions of the autograd engine.
This is, during model inference, moreno_gradefficient, and can further reduce memory usage and improve inference speed.
Function Definition
torch.inference_mode(mode=True)
Parameters:
mode(bool, optional): IfTrue(default value), enable inference mode; ifFalse, exit inference mode. Defaults toTrue。
Return Value:
- Returns a context manager in which gradients and autograd are disabled.
Usage Examples
Example 1: Basic Usage
Example
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
# In the inference_mode context
with torch.inference_mode():
y = x * 2
print("In inference_mode:", y.requires_grad)
# In the no_grad context
with torch.no_grad():
z = x * 2
print("In no_grad:", z.requires_grad)
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
# In the inference_mode context
with torch.inference_mode():
y = x * 2
print("In inference_mode:", y.requires_grad)
# In the no_grad context
with torch.no_grad():
z = x * 2
print("In no_grad:", z.requires_grad)
The output is:
在 inference_mode 中: False 在 no_grad 中: False
Example 2: Comparing no_grad and inference_mode
Example
import torch
# Create tensor
x = torch.randn(100, 100)
# In inference_mode
with torch.inference_mode():
# Performed a lot of computation
for _ in range(10):
x = torch.mm(x, x)
# Even after computation is complete, tensors in the context cannot be used for backpropagation
result = x.sum()
# Check whether it can be converted to a tensor that requires gradients
print("In inference_mode:", result.is_leaf)
# Do the same computation in no_grad
x2 = torch.randn(100, 100)
with torch.no_grad():
for _ in range(10):
x2 = torch.mm(x2, x2)
result2 = x2.sum()
print("In no_grad:", result2.is_leaf)
# Create tensor
x = torch.randn(100, 100)
# In inference_mode
with torch.inference_mode():
# Performed a lot of computation
for _ in range(10):
x = torch.mm(x, x)
# Even after computation is complete, tensors in the context cannot be used for backpropagation
result = x.sum()
# Check whether it can be converted to a tensor that requires gradients
print("In inference_mode:", result.is_leaf)
# Do the same computation in no_grad
x2 = torch.randn(100, 100)
with torch.no_grad():
for _ in range(10):
x2 = torch.mm(x2, x2)
result2 = x2.sum()
print("In no_grad:", result2.is_leaf)
The output is:
在 inference_mode 中: False 在 no_grad 中: True
Example 3: Model Inference
Example
import torch
import torch.nn as nn
# Define a simple model
model = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 5)
)
model.eval()
# Create input data
x = torch.randn(100, 10)
# Use inference_mode for inference
with torch.inference_mode():
output = model(x)
print("Output shape:", output.shape)
print("Output requires_grad:", output.requires_grad)
# Can also use the decorator
@torch.inference_mode()
def predict(x):
return model(x)
result = predict(x)
print("Decorator method - Output shape:", result.shape)
import torch.nn as nn
# Define a simple model
model = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 5)
)
model.eval()
# Create input data
x = torch.randn(100, 10)
# Use inference_mode for inference
with torch.inference_mode():
output = model(x)
print("Output shape:", output.shape)
print("Output requires_grad:", output.requires_grad)
# Can also use the decorator
@torch.inference_mode()
def predict(x):
return model(x)
result = predict(x)
print("Decorator method - Output shape:", result.shape)
The output is:
输出形状: torch.Size([100, 5]) 输出 requires_grad: False 装饰器方式 - 输出形状: torch.Size([100, 5])
Example 4: Memory Optimization Comparison
Example
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(1000, 1000),
nn.ReLU(),
nn.Linear(1000, 1000),
nn.ReLU(),
nn.Linear(1000, 10)
)
# Test memory usage of different modes
x = torch.randn(50, 1000)
print("Without any context manager:")
_ = model(x)
print("\nUsing no_grad:")
with torch.no_grad():
_ = model(x)
print("\nUsing inference_mode:")
with torch.inference_mode():
_ = model(x)
import torch.nn as nn
model = nn.Sequential(
nn.Linear(1000, 1000),
nn.ReLU(),
nn.Linear(1000, 1000),
nn.ReLU(),
nn.Linear(1000, 10)
)
# Test memory usage of different modes
x = torch.randn(50, 1000)
print("Without any context manager:")
_ = model(x)
print("\nUsing no_grad:")
with torch.no_grad():
_ = model(x)
print("\nUsing inference_mode:")
with torch.inference_mode():
_ = model(x)
Usinginference_modecan further optimize memory because it completely disables the autograd engine.
Related Functions
torch.no_grad(): Disables gradient computation but still retains some autograd functionality.torch.enable_grad(): Enables gradient computation.torch.is_inference_mode_enabled(): Checks whether inference mode is enabled.
Notes
inference_modeCompareno_gradstricter, disables more functionality.- In
inference_modeTensors created in it are marked as non-leaf nodes and cannot be used for backpropagation. - It is recommended to use during model inference and evaluation
inference_modeto obtain the best performance. inference_modeCannot beno_gradnested.
Other Extensions