You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改PyTorch的loss.backward()以处理计算中的np.nan问题

类型报错解决方法

首先处理RuntimeError: Found dtype Double but expected Float报错,二选一即可:

  • 将所有输入张量转为float32,输入模型前调用.float()方法
  • 将模型整体转为float64,适配double类型输入,调用model = model.double()

忽略nan值训练实现

不需要修改loss.backward()或者自定义反向传播逻辑,仅需要在计算损失时屏蔽nan对应的位置即可:

  1. 预处理输入张量,将输入中的nan填充为任意占位值(比如0):
# x为输入张量
x = torch.nan_to_num(x, nan=0.0)
  1. 计算损失时生成掩码,仅对标签中存在有效值的位置计算损失:
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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.26 14:06:04