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

jax.nn.softmax中lax.stop_gradient的作用是什么?

JAX Softmax中stop_gradient的必要性解析

官方Softmax实现

JAX官方的jax.nn.softmax实现如下:

def softmax(x: Array,
            axis: Optional[Union[int, Tuple[int, ...]]] = -1,
            where: Optional[Array] = None,
            initial: Optional[Array] = None) -> Array:
  x_max = jnp.max(x, axis, where=where, initial=initial, keepdims=True)
  unnormalized = jnp.exp(x - lax.stop_gradient(x_max))
  return unnormalized / jnp.sum(unnormalized, axis, where=where, keepdims=True)

疑问:stop_gradient(x_max)好像没影响?

我测试了几种Softmax实现,包括不加stop_gradient的稳定版,发现不管是前向计算结果还是梯度结果,和官方实现完全一致:

测试代码与验证结果

import jax
import jax.numpy as jnp

def softmax_unstable(x):
    return jnp.exp(x) / jnp.sum(jnp.exp(x))

def softmax_stable(x):
    x = x - jnp.max(x)
    return jnp.exp(x) / jnp.sum(jnp.exp(x))

def softmax_stop_gradient(x):
    x = x - jax.lax.stop_gradient(jnp.max(x))
    return jnp.exp(x) / jnp.sum(jnp.exp(x))

# 生成测试输入
x = jax.random.normal(jax.random.PRNGKey(123), (100,))

# 验证前向计算结果一致
a = softmax_unstable(x)
b = softmax_stable(x)
c = softmax_stop_gradient(x)
d = jax.nn.softmax(x)
assert jnp.allclose(a, b) and jnp.allclose(b, c) and jnp.allclose(c, d)

# 验证单次Softmax的梯度一致
a = jax.grad(lambda x: -jnp.log(softmax_unstable(x))[2])(x)
b = jax.grad(lambda x: -jnp.log(softmax_stable(x))[2])(x)
c = jax.grad(lambda x: -jnp.log(softmax_stop_gradient(x))[2])(x)
d = jax.grad(lambda x: -jnp.log(jax.nn.softmax(x))[2])(x)
assert jnp.allclose(a, b) and jnp.allclose(b, c) and jnp.allclose(c, d)

# 验证嵌套Softmax的梯度一致
a = jax.grad(lambda x: -jnp.log(softmax_unstable(softmax_unstable(x)))[2])(x)
b = jax.grad(lambda x: -jnp.log(softmax_stable(softmax_stable(x)))[2])(x)
c = jax.grad(lambda x: -jnp.log(softmax_stop_gradient(softmax_stop_gradient(x)))[2])(x)
d = jax.grad(lambda x: -jnp.log(jax.nn.softmax(jax.nn.softmax(x)))[2])(x)
assert jnp.allclose(a, b) and jnp.allclose(b, c) and jnp.allclose(c, d)

所有测试都通过了,那这个stop_gradient到底有什么用?

实际作用:优化反向传播

从数学上看,有无stop_gradient确实不影响最终的梯度结果,但它能带来两个关键好处:

  • 减少反向传播的计算量与内存占用
    没有stop_gradient时,JAX会追踪jnp.max(x)的梯度计算逻辑,但从导数推导可知,这部分梯度最终会和其他项完全抵消,属于无用计算。加了stop_gradient后,直接切断了x_max到输入x的梯度传递路径,跳过这部分冗余计算,让反向传播更快、内存占用更低。

  • 避免极端场景的数值不稳定
    如果输入x中存在极大值元素,计算x_max的梯度时可能出现数值异常(比如梯度爆炸或NaN)。虽然这部分异常会被后续计算抵消,但stop_gradient能从根源上避免这种潜在问题,让反向传播过程更鲁棒。

简言之,stop_gradient在这里不改变最终结果,但能让反向传播更高效、更稳定。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 08:10:29