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

训练采用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 10:09:04