代码片段记录16

1️⃣[Pytorch检查tensor nan]

1
2
3
4
5
6
7
8
9
10
# 法一,基于nan!=nan 
>>> x = torch.tensor([1, 2, np.nan])
tensor([ 1., 2., nan.])
>>> x != x
tensor([ 0, 0, 1], dtype=torch.uint8)

# 法二,torch.isnan(x)

>>> torch.isnan(x)
tensor([ 0, 0, 1], dtype=torch.uint8)