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
相关产品推荐
相关产品推荐

