训练采用BCE损失的Siamese网络触发IndexError索引越界如何解决
错误原因定位
报错出现的核心原因是y_true.long()返回了非法的索引值-9223372036854775808,这个值是PyTorch中NaN值转换为int64类型的默认结果,触发场景如下:
- 当某一个batch内所有的
y_true值都为0时,torch.max(y_true)返回值为0,你对maximum_y_true做nan_to_num处理后仍然为0 - 后续执行
y_true = y_true / maximum_y_true时出现除以0的操作,导致y_true所有值变为NaN - NaN转换为long类型后得到非法索引,访问第二维大小为2的
y_true_2时触发越界错误
解决方案
你可以选择以下任意一种方案解决问题:
方案1:修复归一化逻辑,规避非法索引
在计算y_true归一化的步骤添加保护逻辑,避免除以0,同时保证y_true转long后是合法的0/1索引:
maximum_y_true = torch.max(y_true) maximum_y_true = torch.nan_to_num(maximum_y_true) # 新增:最大值为0时默认用1做分母,避免除以0 if maximum_y_true < 1e-8: maximum_y_true = torch.tensor(1.0) y_true = y_true / maximum_y_true # 新增:截断y_true到0~1范围,四舍五入后转long,保证索引合法 y_true = torch.clamp(y_true.round(), min=0, max=1).long()
方案2:简化BCE损失计算逻辑,无需构造one-hot标签
你当前的BCE损失使用方式可以直接简化为1维输入,完全规避索引赋值的风险,修改成本更低:
删除以下几行代码:
input_dy = torch.empty(dy.size(0), 2) input_dy[:, 0] = 1 - dy input_dy[:, 1] = dy y_true_2 = torch.zeros(dy.size(0), 2) y_true_2[range(y_true_2.shape[0]), y_true.long()] = 1 m = nn.Sigmoid() loss = loss_fn(m(input_dy), y_true_2)
替换为:
m = nn.Sigmoid() loss = loss_fn(m(dy), y_true.float()) # 后续指标计算如果需要标签可以直接用1维的y_true,对应修改指标的输入逻辑即可
内容的提问来源于stack exchange,提问作者aarya
相关产品推荐
相关产品推荐

