PyTorch torch.nansum Function


Pytorch torch 参考手册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))

The output result is:

非NaN值之和: tensor(11.)
dim=0 非NaN值之和: tensor([5., 1., 5.])
dim=1 非NaN值之和: tensor([3., 8.])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions