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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 08:24:26