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

为何jnp.round与np.round计算结果存在差异?

JAX与NumPy round函数浮点精度差异问题解析

问题复现

以下代码可复现jnp.round与np.round的结果差异:

import jax.numpy as jnp
import numpy as np
from numpy.testing import assert_array_equal

a = jnp.array([-0.78073686, -0.7908204 ,  2.174842])
b = np.array(a, dtype='float32')
assert_array_equal(a, b)  # 原始数组完全相等

print(a.round(2), a.dtype)
print(b.round(2), b.dtype)

输出结果:

[-0.78       -0.78999996  2.1699998 ] float32
[-0.78 -0.79  2.17] float32

执行精确相等断言时触发报错:

assert_array_equal(a.round(2), b.round(2))

报错信息:

AssertionError: 
Arrays are not equal

Mismatched elements: 2 / 3 (66.7%)
Max absolute difference: 2.3841858e-07
Max relative difference: 1.0987031e-07
 x: array([-0.78, -0.79,  2.17], dtype=float32)
 y: array([-0.78, -0.79,  2.17], dtype=float32)

注:直接定义b = np.array([-0.78073686, -0.7908204 , 2.174842], dtype='float32')也会得到相同结果,排除数组转换环节的问题。

原因分析

  1. float32精度限制:float32类型仅能保留约6-7位有效数字,部分十进制小数无法被二进制浮点格式精确表示,舍入操作会触发微小的截断误差。
  2. 底层实现差异:JAX为适配GPU/TPU硬件加速,round函数可能采用了硬件原生舍入指令;而NumPy的实现偏向通用CPU场景,二者在边缘浮点值的舍入处理细节上存在差异,最终导致结果的二进制表示出现可忽略的偏差。

解决方案

浮点运算场景下不应追求精确相等,应使用允许误差范围的断言函数:

  • 用np.testing.assert_allclose替代assert_array_equal,设置适配float32精度的相对误差(rtol)和绝对误差(atol)阈值。

示例代码:

np.testing.assert_allclose(a.round(2), b.round(2), rtol=1e-6, atol=1e-7)

该断言会自动忽略浮点运算中符合精度范围的微小差异,避免误报。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 13:06:01