Jax中含argmax的损失函数梯度全为零如何解决
问题原因
jnp.argmax是输出离散类别索引的分段常数算子,数学性质决定了它除了两个类别得分完全相等的跳变点外,其余位置的梯度恒为0:当模型参数发生微小变化时,只要没有大到改变最大logit的位置,argmax的输出就不会发生任何变化,反向传播时梯度无法穿过argmax算子回传到模型参数,最终得到全零梯度。
当前代码的核心问题是训练阶段直接将连续可导的模型logit输出通过argmax转换为离散的硬预测值,完全切断了损失和模型参数之间的梯度传播链路。
解决方案
优先选择训练阶段使用连续可导的算子计算损失,仅在推理阶段使用argmax获取最终类别结果,这是分类任务训练的标准做法:
- 多分类/二分类任务中,将模型输出的logit通过softmax转换为连续的类别概率分布,再搭配交叉熵损失计算梯度,整个计算链路全程可导,不会出现梯度消失的问题。argmax仅在推理阶段用来把概率转回离散类别标签即可。
修改后的损失函数示例:
def loss_fn(parameters, x: chex.Array, y: chex.Array): y_hat = apply(parameters, x) # 训练阶段用连续可导的概率输出计算损失,不插入argmax log_probs = jax.nn.log_softmax(y_hat, axis=1) # 交叉熵损失:y为0/1类别标签时,取对应类别的负对数似然求均值 return -jnp.take_along_axis( log_probs, y.astype(jnp.int32).reshape(-1, 1), axis=1 ).mean()
如果场景必须在训练过程中保留argmax的硬输出特性,可以使用带梯度近似的可微argmax替代原生算子:
- 直通估计器(Straight-Through Estimator):前向传播时正常计算argmax的离散输出,反向传播时直接绕过argmax将梯度回传,实现简单,梯度存在一定估计偏差。
- Soft-Argmax:对logit加温度系数做softmax后,与类别索引值做加权求和,温度越低输出越接近真实argmax的结果,全程可导,温度参数需要根据场景调优。
直通估计器的简单实现参考:
import jax @jax.custom_vjp def ste_argmax(x, axis=-1): return jnp.argmax(x, axis=axis) def ste_argmax_fwd(x, axis=-1): return ste_argmax(x, axis), (x, axis) def ste_argmax_bwd(res, g): # 反向传播时直接将梯度传回输入,跳过argmax的零梯度 return g, None ste_argmax.defvjp(ste_argmax_fwd, ste_argmax_bwd)
注意:可微近似argmax存在梯度估计偏差,除非场景有特殊要求,优先使用softmax+交叉熵的标准分类训练流程,不要在损失计算链路中插入原生argmax算子。
内容的提问来源于stack exchange,提问作者Light
相关产品推荐
相关产品推荐

