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, returns
True; 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())
# 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())
"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)
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)
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
AI thinking...