TensorFlow中对应PyTorch autograd.detect_anomaly的NaN调试工具是什么?
TensorFlow NaN损失定位工具及调试方案
原生等效工具
TensorFlow官方提供了和torch.autograd.detect_anomaly功能完全匹配的NaN/Inf检测能力,你只需要在训练代码启动的最开头调用一次tf.debugging.enable_check_numerics()即可开启全局校验。运行过程中任意算子输出NaN或Inf时,会立刻抛出异常,同时输出触发异常的算子位置、张量名称、完整调用栈信息,可直接定位到问题代码。
如果不需要全局检测,也可以针对可疑的计算节点单独添加校验,在对应计算逻辑后追加:
output = tf.debugging.check_numerics(output, "自定义可疑节点标识")
当该节点输出异常时,会直接抛出携带你设置的标识的报错,更快缩小排查范围。
补充调试方案
- 区分前向/反向传播问题:先在损失计算完成后添加
check_numerics校验,如果损失阶段就触发异常,说明问题出在前向传播(损失计算、算子逻辑、输入处理),否则属于反向传播梯度计算问题 - 临时关闭图编译优化:运行前设置
tf.config.run_functions_eagerly(True)关闭tf.function的图编译优化,报错时会输出逐行对应的Python代码调用栈,避免图模式栈信息难以和业务代码对应 - 混合精度场景专项排查:如果开启了混合精度训练,先关闭混合精度运行测试,如果不再出现NaN,说明是loss scale设置过大导致梯度上溢,更换动态loss scale或者调小固定loss scale即可解决
- 输入合法性校验:即使你已经设置了权重约束,若某批输入存在极端值、标签非法值,也可能触发
log(0)、除零等操作产生异常值,可添加输入/标签合法性校验过滤无效样本
内容的提问来源于stack exchange,提问作者Achintha Ihalage
相关产品推荐
相关产品推荐

