PyTorch torch.inference_mode Function


Pytorch torch 参考手册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)

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)

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)

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)

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.
  • Ininference_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 evaluationinference_modeto obtain the best performance.
  • inference_modeCannot beno_gradnested.

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions