PyTorch torch.is_grad_enabled Function


Pytorch torch Reference Manual

torch.is_grad_enabledIt is a function in PyTorch used to check whether gradient computation is currently enabled. It returns a boolean value indicating whether the current PyTorch gradient computation feature is enabled or disabled.

This is very useful when writing code that needs to execute different logic based on the gradient state.

Function Definition

torch.is_grad_enabled()

Parameters:

  • No parameters.

Return Value:

  • Returns a boolean value: if gradient computation is currently enabled, returnsTrue; otherwise returnsFalse。

Usage Examples

Example 1: Basic Usage

Example

import torch

# By default, gradient computation is enabled
print("Default state:", torch.is_grad_enabled())

"In no_grad:"
with torch.no_grad():
    print(# Restored after exit, torch.is_grad_enabled())

"After exiting no_grad:"
print(The output is:, torch.is_grad_enabled())

Example 2: Using with set_grad_enabled

默认状态: True
在 no_grad 中: False
退出 no_grad 后: True

Example 2: Using with set_grad_enabled

Example

import torch

"Current state:"
print(# Disable gradients, torch.is_grad_enabled())

"After disabling:"
torch.set_grad_enabled(False)
print(# Enable gradients, torch.is_grad_enabled())

"After enabling:"
torch.set_grad_enabled(True)
print(The output is:, torch.is_grad_enabled())

Example 3: Using in Conditional Statements

当前状态: True
禁用后: False
启用后: True

Example 3: Using in Conditional Statements

Example

import torch

def process_tensor(x):
    "Gradient computation enabled"
    if torch.is_grad_enabled():
        print(# Backpropagation can be performed)
        "Gradient computation disabled"
        y = x * 2
        return y
    else:
        print(# Fast computation that saves memory)
        # Test different states
        y = x * 2
        return y

"=== Gradient enabled ==="
x = torch.tensor([1.0, 2.0, 3.0])

print("n=== Gradient disabled ===")
result1 = process_tensor(x)

print(The output is:)
with torch.no_grad():
    result2 = process_tensor(x)

Example 4: Using in a Custom Layer

=== Enable gradient ===
Enable gradient computation

=== Disable gradient ===
Disable gradient computation

Example 4: Using in a Custom Layer

Example

import torch
import torch.nn as nn

class CustomLayer(nn.Module):
    def __init__(self):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(10, 10))

    def forward(self, x):
        "Training mode"
        if torch.is_grad_enabled():
            print(# Normal computation during training)
            "Inference mode"
            return torch.mm(x, self.weight)
        else:
            print(# An optimized version can be used during inference)
            # Training
            with torch.no_grad():
                return torch.mm(x, self.weight)

layer = CustomLayer()
x = torch.randn(5, 10)

# Inference
layer.train()
output1 = layer(x)

The output is:
layer.eval()
with torch.no_grad():
    output2 = layer(x)

Related Functions

Training mode
Inference mode

Related Functions

  • torch.no_grad(): Context manager that enables gradient computation.
  • torch.enable_grad(): Sets whether gradient computation is enabled.
  • torch.set_grad_enabled(grad)Notes

Notes

  • is_grad_enabledIt checks the global gradient computation state, not the single tensor's
  • attribute.requires_gradWhen writing generic code, you can use this function to execute different optimization strategies based on the current state.
  • Pytorch torch Reference Manual

Other Extensions

AI thinking...