为何jnp.einsum与手动循环计算的内积结果存在差异?
问题描述
定义两个JAX矩阵:
import jax import jax.numpy as jnp a = jax.random.normal(jax.random.PRNGKey(0), shape=(64,16), dtype=jnp.float32) b = jax.random.normal(jax.random.PRNGKey(1), shape=(64,16), dtype=jnp.float32)
用两种方式计算最后维度的内积:
- 方式1:
jnp.einsum实现
inner_prod1 = jnp.einsum('i d, j d -> i j', a, b)
- 方式2:循环调用
jnp.dot手动赋值
inner_prod2 = jnp.zeros((64,64)) for i1 in range(64): for i2 in range(64): inner_prod2 = inner_prod2.at[i1, i2].set(jnp.dot(a[i1], b[i2]))
执行后两者差值的最大值为0.03830552,数学上等价但结果差异明显,请问原因是什么?
原因分析
- 浮点数运算的固有精度限制:
jnp.float32是单精度浮点数,仅能保留约6-7位十进制有效数字。浮点数加法不满足严格的结合律,不同的计算顺序会导致舍入误差的累积结果不同。 - 两种实现的计算路径差异:
jnp.einsum本质等价于矩阵乘法a @ b.T,JAX会调用底层优化的BLAS库(如MKL、cuBLAS)执行计算。这类库会采用向量化、分块并行等高效策略,甚至可能在累加过程中使用更高精度的中间值(如临时用float64累加再转回float32),或者按特定的块顺序完成乘积求和,误差累积的方式与手动循环完全不同。- 手动循环的方式是逐个计算16维向量的点积:对每一对
a[i1]和b[i2],先计算对应元素相乘,再将16个乘积结果逐个累加为单精度浮点数。这种逐元素累加的顺序,与BLAS库优化后的矩阵乘法累加顺序不一致,最终导致舍入误差的累积结果出现差异。
- JAX的不可变数组特性:手动循环中每次调用
at[i1, i2].set都会创建新的数组副本,虽然不影响数学逻辑,但计算过程的内存访问和操作顺序与einsum的批量优化路径完全不同,进一步放大了浮点数误差的表现。
内容的提问来源于stack exchange,提问作者qchuG
相关产品推荐
相关产品推荐

