PyTorch torch.nansum Function
Pytorch torch Reference Manual
torch.nansumIt is a function in PyTorch for returning the sum of non-NaN values of a tensor. Unlike sum, it ignores all NaN values.
Function Definition
torch.nansum(input, dim, keepdim=False)
Usage Example
Example
import torch
x = torch.tensor([1.0, float('nan'), 3.0, 2.0, 5.0])
# Return the sum of non-NaN values
print("Sum of non-NaN values:", torch.nansum(x))
# Sum of non-NaN values along dim=0
y = torch.tensor([[1.0, float('nan'), 2.0], [4.0, 1.0, 3.0]])
print("dim=0 sum of non-NaN values:", torch.nansum(y, dim=0))
print("dim=1 sum of non-NaN values:", torch.nansum(y, dim=1))
x = torch.tensor([1.0, float('nan'), 3.0, 2.0, 5.0])
# Return the sum of non-NaN values
print("Sum of non-NaN values:", torch.nansum(x))
# Sum of non-NaN values along dim=0
y = torch.tensor([[1.0, float('nan'), 2.0], [4.0, 1.0, 3.0]])
print("dim=0 sum of non-NaN values:", torch.nansum(y, dim=0))
print("dim=1 sum of non-NaN values:", torch.nansum(y, dim=1))
The output result is:
非NaN值之和: tensor(11.) dim=0 非NaN值之和: tensor([5., 1., 5.]) dim=1 非NaN值之和: tensor([3., 8.])
Other Extensions