PyTorch torch.is_inference_mode_enabled Function


Pytorch torch 参考手册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, returnsTrue; 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())

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())

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)

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)

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 detectinference_mode, notno_grad。
  • Ininference_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.

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions