PyTorch torch.no_grad Function


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

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)

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)

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)

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 usetorch.no_grad()to save memory and improve speed.
  • Ifno_gradvariables created inside the block need to be used outside the block, they need to be manually copied out.
  • andmodel.eval()Use in combination for best results.

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions