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

PyTorch张量溢出与NaN问题排查及解决请求

问题排查与解决方案:PyTorch大张量溢出导致NaN及CUDA断言错误

问题根源分析

  1. 张量溢出原因:你看到的-9223372036854775808是torch.int64类型的最小值,说明存储索引的张量使用了过小的数据类型(如torch.int32),或者索引数值本身超出了int64范围,导致整数溢出后被 wrap 到类型最小值。
  2. NaN与CUDA断言关联:溢出后的无效值(负数最小值)在neural-astar的train_maps.py计算逻辑中,被参与浮点运算(如除法、开方、归一化)时会转换为-inf,进一步计算后产生NaN;CUDA设备端断言错误则是因为NaN/inf参与了需要有效数值的操作(如张量索引、scatter_等),触发了CUDA内核的合法性检查。

具体修复步骤

1. 强制使用足够大的数据类型存储索引

确保所有存储业务索引的张量使用torch.int64类型(PyTorch支持的最大整数类型),避免溢出:

# 替换原张量创建代码,显式指定dtype
idx_tensor = torch.tensor(raw_index_data, dtype=torch.int64, device=device)

如果索引数值超出int64范围(最大值为9223372036854775807),需对索引做映射处理:

  • 采用相对索引:将所有索引减去最小值,把数值范围压缩到int64有效区间内
  • 采用哈希映射:用字典将超大索引映射到较小的连续整数,保留业务关联的同时避免溢出

2. 修复train_maps.py中的idx计算逻辑

定位idx变量的生成代码,添加无效值过滤与数值校验:

# 假设overflowed_idx是存储索引的输入张量
# 替换溢出的无效值为合法最小值(根据业务逻辑调整)
valid_idx = torch.where(
    overflowed_idx == -9223372036854775808,
    torch.tensor(1, dtype=torch.int64, device=device),
    overflowed_idx
)
# 转换为浮点类型时使用高精度的float64,避免转换异常
float_idx = valid_idx.to(torch.float64)
# 执行原计算逻辑前,断言无无效值
assert not torch.isinf(float_idx).any(), "Invalid inf values detected in index tensor"
# 后续计算
idx = your_original_calculation(float_idx)
# 最后检查idx是否产生NaN
assert not torch.isnan(idx).any(), "NaN values generated in idx variable"

3. 排查CUDA断言触发点

CUDA断言错误通常会给出具体的内核调用位置,根据报错日志定位到对应代码:

  • 如果是索引操作触发断言,确保用于索引的张量无NaN/inf,且数值在合法的索引范围内
  • 如果是归一化、损失计算等操作触发,添加torch.clamp()限制数值范围:
# 限制数值在有效区间内,避免NaN/inf
normalized_idx = torch.clamp(idx, min=1e-8, max=1e8)

数据集层面的预防处理

检查数据集的索引字段:

  • 确认所有索引数值的范围,若存在超出int64的情况,提前在数据预处理阶段做映射转换
  • 添加数据校验逻辑,过滤或修正无效的超大索引值

内容的提问来源于stack exchange,提问作者SamuelMastrelli

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 08:56:13