PyTorch torch.enable_grad Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.enable_gradIt is a context manager in PyTorch used to enable gradient computation. It istorch.no_gradthe opposite, and is used to explicitly enable gradient computation in code blocks that require gradients.

This is very useful when you temporarily need to train a model in inference code, or when you need to enable gradient computation in a local region.

Function Definition

torch.enable_grad()

Parameters:

  • No parameters. This is a context manager.

Return Value:

  • Returns a context manager that enables gradient computation in that context.

Usage Examples

Example 1: Basic Usage

Example

import torch

# By default, in a no_grad context
with torch.no_grad():
    x = torch.tensor([1.0, 2.0, 3.0])
    print("In no_grad:", x.requires_grad)

    # Use enable_grad to temporarily enable gradients
    with torch.enable_grad():
        y = x * 2
        print("In enable_grad:", y.requires_grad)

    # After exiting, restore no_grad state
    z = x * 2
    print("After exiting:", z.requires_grad)

The output is:

在 no_grad 中: False
在 enable_grad 中: True
退出后: False

Example 2: Mixing Training and Inference

Example

import torch
import torch.nn as nn

model = nn.Linear(10, 2)

# Inference mode
with torch.no_grad():
    x = torch.randn(5, 10)
    output1 = model(x)
    print("Inference output:", output1.shape)

    # If you need to temporarily train some parameters during inference
    with torch.enable_grad():
        # Create a tensor that requires gradients for computation
        temp_weight = torch.randn(10, 10, requires_grad=True)
        temp_output = torch.mm(x, temp_weight)
        print("Temporarily enable gradients:", temp_output.requires_grad)

The output is:

推理输出: torch.Size([5, 2])
临时启用梯度: True

Example 3: Using as a Decorator

Example

import torch

@torch.enable_grad()
def train_step(x, y):
    """Simulate training steps"""
    # This function will enable gradient computation internally
    loss = (x - y).sum()
    return loss

# Call it in a no_grad context
with torch.no_grad():
    x = torch.tensor([1.0, 2.0, 3.0])
    y = torch.tensor([0.0, 0.0, 0.0])
    # The decorator ensures gradient computation is enabled
    loss = train_step(x, y)
    print("Loss:", loss)
    print("Requires gradient:", loss.requires_grad)

The output is:

Loss: tensor(6.)
需要梯度: True

Related Functions

  • torch.no_grad(): Disables gradient computation.
  • torch.set_grad_enabled(grad): Enables or disables gradient computation according to parameters.
  • torch.is_grad_enabled(): Checks whether gradient computation is currently enabled.

Notes

  • enable_gradMainly used tono_gradtemporarily enable gradient computation in a context.
  • If used in an environment where gradients are globally enabledenable_grad, it will have no effect.
  • It is recommended to useenable_gradthe decorator to ensure that gradients are always enabled inside the function.

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions