如何在含NaN的PyTorch二维张量中按列求最小值并忽略NaN
PyTorch:忽略NaN计算二维张量各列最小值
直接用PyTorch内置的torch.nanmin函数就能解决,它会自动忽略张量中的NaN值,计算有效元素的最小值:
import torch data = torch.tensor([[ 0., 1., float('nan'), 3.],[ 4., 5., 6., 7.], [ 8., 9., 10., 11.]]) result = torch.nanmin(data, dim=0)[0] print(result) # 输出:tensor([0., 1., 6., 3.])
为什么原来的torch.min不行?
torch.min在遇到NaN时会直接返回NaN(因为NaN和任何数值比较的结果都是NaN),所以第三列会得到NaN,不符合你的需求。
低版本PyTorch兼容方案(仅作参考)
如果你的PyTorch版本低于1.8.0(torch.nanmin在1.8.0版本新增),可以手动把NaN替换为正无穷,再用torch.min计算:
data_filled = data.clone() # 将NaN替换为正无穷,求最小值时不会影响有效元素的结果 data_filled[torch.isnan(data_filled)] = float('inf') result = torch.min(data_filled, dim=0)[0] print(result) # 输出同样符合预期:tensor([0., 1., 6., 3.])
内容的提问来源于stack exchange,提问作者tamnva
相关产品推荐
相关产品推荐

