PyTorch torch.is_inference_mode_enabled Function
PyTorch torch Reference Manual
torch.is_inference_mode_enabledis a function in PyTorch used to check whether inference mode is currently enabled. It returns whether the current context is ininference_modea boolean value of the context.
This is useful when writing code that needs to execute different logic based on the inference mode state.
Function Definition
torch.is_inference_mode_enabled()
Parameters:
- No parameters.
Return Value:
- Returns a boolean: if inference mode is currently enabled, returns
True; otherwise returnsFalse。
Usage Examples
Example 1: Basic Usage
Example
import torch
# By default, inference mode is disabled
print("Default state:", torch.is_inference_mode_enabled())
# Inside the inference_mode context
with torch.inference_mode():
print("Inside inference_mode:", torch.is_inference_mode_enabled())
# Restored after exiting
print("After exiting:", torch.is_inference_mode_enabled())
# By default, inference mode is disabled
print("Default state:", torch.is_inference_mode_enabled())
# Inside the inference_mode context
with torch.inference_mode():
print("Inside inference_mode:", torch.is_inference_mode_enabled())
# Restored after exiting
print("After exiting:", torch.is_inference_mode_enabled())
The output result is:
默认状态: False 在 inference_mode 中: True 退出后: False
Example 2: Comparing is_grad_enabled and is_inference_mode_enabled
Example
import torch
# Inside no_grad
with torch.no_grad():
print("Inside no_grad:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
# Inside inference_mode
with torch.inference_mode():
print("Inside inference_mode:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
# Default state
print("Default state:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
# Inside no_grad
with torch.no_grad():
print("Inside no_grad:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
# Inside inference_mode
with torch.inference_mode():
print("Inside inference_mode:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
# Default state
print("Default state:")
print(" is_grad_enabled:", torch.is_grad_enabled())
print(" is_inference_mode_enabled:", torch.is_inference_mode_enabled())
The output result is:
在 no_grad 中: is_grad_enabled: False is_inference_mode_enabled: False 在 inference_mode 中: is_grad_enabled: False is_inference_mode_enabled: True 默认状态: is_grad_enabled: True is_inference_mode_enabled: False
Example 3: Using in Conditional Judgments
Example
import torch
import torch.nn as nn
def forward_pass(x, model):
"""Execute different optimizations based on inference mode"""
if torch.is_inference_mode_enabled():
print("Using inference mode optimization")
# Optimized computation under inference mode
return model(x)
elif torch.is_grad_enabled():
print("Training mode")
return model(x)
else:
print("Eval mode")
return model(x)
model = nn.Linear(10, 5)
x = torch.randn(5, 10)
# Test different states
print("=== Training mode ===")
with torch.enable_grad():
result = forward_pass(x, model)
print("n=== Inference mode ===")
with torch.inference_mode():
result = forward_pass(x, model)
print("n=== Eval mode ===")
with torch.no_grad():
result = forward_pass(x, model)
import torch.nn as nn
def forward_pass(x, model):
"""Execute different optimizations based on inference mode"""
if torch.is_inference_mode_enabled():
print("Using inference mode optimization")
# Optimized computation under inference mode
return model(x)
elif torch.is_grad_enabled():
print("Training mode")
return model(x)
else:
print("Eval mode")
return model(x)
model = nn.Linear(10, 5)
x = torch.randn(5, 10)
# Test different states
print("=== Training mode ===")
with torch.enable_grad():
result = forward_pass(x, model)
print("n=== Inference mode ===")
with torch.inference_mode():
result = forward_pass(x, model)
print("n=== Eval mode ===")
with torch.no_grad():
result = forward_pass(x, model)
The output result is:
=== 训练模式 === 训练模式 === 推理模式 === 使用推理模式优化 === eval 模式 === eval 模式
Example 4: Detecting in Custom Modules
Example
import torch
import torch.nn as nn
class OptimizedLayer(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.randn(20, 10))
def forward(self, x):
# Detect inference mode and apply optimization
if torch.is_inference_mode_enabled():
# Inference mode: use more efficient computation
return torch.mm(x, self.weight)
elif torch.is_grad_enabled():
# Training mode: preserve gradient computation
return torch.mm(x, self.weight)
else:
# Eval mode: gradients not needed but optimization can be used
with torch.no_grad():
return torch.mm(x, self.weight)
layer = OptimizedLayer()
x = torch.randn(5, 10)
print("=== Training mode ===")
with torch.enable_grad():
_ = layer(x)
print("n=== Inference mode ===")
with torch.inference_mode():
_ = layer(x)
print("n=== Eval mode ===")
with torch.no_grad():
_ = layer(x)
import torch.nn as nn
class OptimizedLayer(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.randn(20, 10))
def forward(self, x):
# Detect inference mode and apply optimization
if torch.is_inference_mode_enabled():
# Inference mode: use more efficient computation
return torch.mm(x, self.weight)
elif torch.is_grad_enabled():
# Training mode: preserve gradient computation
return torch.mm(x, self.weight)
else:
# Eval mode: gradients not needed but optimization can be used
with torch.no_grad():
return torch.mm(x, self.weight)
layer = OptimizedLayer()
x = torch.randn(5, 10)
print("=== Training mode ===")
with torch.enable_grad():
_ = layer(x)
print("n=== Inference mode ===")
with torch.inference_mode():
_ = layer(x)
print("n=== Eval mode ===")
with torch.no_grad():
_ = layer(x)
The output result is:
=== 训练模式 === 训练模式 === 推理模式 === 推理模式:使用更高效的计算方式 === eval 模式 === eval 模式
Related Functions
torch.inference_mode(): Context manager that enables inference mode.torch.no_grad(): Context manager that disables gradient computation.torch.is_grad_enabled(): Checks whether gradient computation is enabled.
Notes
is_inference_mode_enabledIt is a read-only function and does not change any state.- It is specifically used to detect
inference_mode, notno_grad。 - In
inference_modeWhen inis_grad_enabledalso returnsFalse, but the converse is not true. - When writing generic code, you can use this function to distinguish different running modes.
Other Extensions