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

为何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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 05:22:29