自定义Mish激活函数实现出现Loss=NaN问题求助
自定义Mish激活函数训练出现NaN的问题分析
实验配置
- 模型架构:VGG 16
- 数据集:CIFAR-10
- 训练轮次:100
- 学习率:前80轮0.001,后20轮0.0001
- 优化器:Adam
- 损失函数:categorical cross entropy
- 批次大小:128
- 卷积层数量列表:
num_cnn_list = [2,2,4,4,4] - 通道数列表:
channel_list = [16,128,256,512,512]
问题现象
使用原生softplus实现Mish时训练无异常:
@tf.function def mish(x): return x * K.tanh(K.softplus(x)) # 等价于x * tanh(ln(1+e^x))
但拆分手动计算ln(1+e^x)时,训练初期(第一个轮次未完成)就出现Loss=NaN:
@tf.function def mish_inside(x): return K.log(1 + K.exp(x)) @tf.function def mish(x): return x * K.tanh(mish_inside(x))
降低学习率至0.0001可正常运行,但希望实现具备多场景鲁棒性。
核心原因:数值稳定性差异
理论上softplus(x)与ln(1+e^x)数学等价,但两者的数值实现稳定性有本质区别:
- 手动实现的
K.log(1+K.exp(x))在输入x较大时(比如x>20),K.exp(x)会超出float32的数值范围变成inf,后续K.log(inf)仍为inf,反向传播时exp(x)的梯度就是自身,直接变成inf,导致参数更新时出现NaN。 - TensorFlow原生
softplus内部做了数值稳定优化:当x超过阈值时直接返回x,避免exp(x)溢出,同时保证正向计算和反向传播的梯度都处于合理范围。
解决方案
1. 给手动实现添加数值稳定分支
修改mish_inside函数,加入阈值判断避免溢出:
@tf.function def mish_inside(x): # 当x>20时,exp(x)会溢出,直接返回x(此时ln(1+e^x)≈x) return tf.where(x > 20., x, tf.math.log(1 + tf.math.exp(x))) @tf.function def mish(x): return x * tf.math.tanh(mish_inside(x))
2. 优先使用原生softplus
直接调用K.softplus或tf.nn.softplus,无需手动实现,原生实现已内置所有数值稳定处理,能最大化鲁棒性。
补充说明
降低学习率能暂时解决问题,是因为小学习率限制了参数更新的幅度,压制了梯度爆炸的影响,但这不是根本解决方案,只有保证激活函数的数值稳定性,才能在不同学习率、不同模型规模下稳定运行。
内容的提问来源于stack exchange,提问作者pnpsuM
相关产品推荐
相关产品推荐

