如何修改PyTorch的loss.backward()以处理计算中的np.nan问题
类型报错解决方法
首先处理RuntimeError: Found dtype Double but expected Float报错,二选一即可:
- 将所有输入张量转为float32,输入模型前调用
.float()方法 - 将模型整体转为float64,适配double类型输入,调用
model = model.double()
忽略nan值训练实现
不需要修改loss.backward()或者自定义反向传播逻辑,仅需要在计算损失时屏蔽nan对应的位置即可:
- 预处理输入张量,将输入中的nan填充为任意占位值(比如0):
# x为输入张量 x = torch.nan_to_num(x, nan=0.0)
- 计算损失时生成掩码,仅对标签中存在有效值的位置计算损失:
import torch.nn.functional as F pred = model(x) # 生成掩码:非nan的位置为True,参与损失计算 mask = ~torch.isnan(y) # 仅计算有效位置的MSE损失,此时loss不会出现nan loss = F.mse_loss(pred[mask], y[mask]) # 正常执行反向传播和参数更新即可 loss.backward() optimizer.step()
上述方案通过索引过滤掉了nan对应的计算节点,反向传播时不会涉及nan相关的计算,完全可以正常训练。训练完成后模型即可实现输入带nan的数组直接输出补全后的完整数组。
内容的提问来源于stack exchange,提问作者Galen BlueTalon
相关产品推荐
相关产品推荐

