PyTorch torch.nanmean Function
PyTorch torch Reference Manual
torch.nanmeanIt is a function in PyTorch used to calculate the average value while ignoring NaN values.
Function Definition
torch.nanmean(input, dim, keepdim=False, out=None)
Usage Example
Example
import torch
# Create a tensor containing NaN
x = torch.tensor([1.0, 2.0, float('nan'), 4.0, 5.0])
# Compute the mean of non-NaN values
mean = torch.nanmean(x)
print("Non-NaN mean:", mean)
# Create a tensor containing NaN
x = torch.tensor([1.0, 2.0, float('nan'), 4.0, 5.0])
# Compute the mean of non-NaN values
mean = torch.nanmean(x)
print("Non-NaN mean:", mean)
The output result is:
非 NaN 均值: tensor(3.)
Other Extensions