如何过滤PyTorch一维张量中的np.nan值,得到尺寸为4的张量?
过滤PyTorch张量中的NaN值
嘿,这事儿其实挺容易搞定的,PyTorch提供了直接检测和过滤NaN的工具,给你两种常用的方案:
方法一:布尔索引(最直观)
你可以先用torch.isnan()生成一个布尔掩码,标记出所有NaN的位置,然后通过取反掩码来筛选出非NaN的元素:
import numpy as np import torch my_list = [0, 1, 2, np.nan, np.nan, 4] tensor = torch.Tensor(my_list) # 生成非NaN的掩码,再索引张量 filtered_tensor = tensor[~torch.isnan(tensor)]
这里torch.isnan(tensor)会返回一个和原张量同形状的布尔张量,~符号用来取反(把True和False互换),最后用这个取反后的掩码去索引原张量,就直接得到了不含NaN的新张量。
方法二:使用torch.masked_select()
如果你想更明确地用PyTorch的官方选择函数,torch.masked_select()专门干这个事儿,用法和布尔索引类似:
filtered_tensor = torch.masked_select(tensor, ~torch.isnan(tensor))
这个函数会根据传入的掩码(这里是非NaN的位置),返回所有符合条件的元素,结果同样是尺寸为4的一维张量。
验证结果
你可以打印一下filtered_tensor,会得到:
tensor([0., 1., 2., 4.])
查看形状的话,filtered_tensor.shape会输出torch.Size([4]),完全符合你的需求。
内容的提问来源于stack exchange,提问作者Mathias Byskov
相关产品推荐
相关产品推荐

