为何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')也会得到相同结果,排除数组转换环节的问题。
原因分析
- float32精度限制:float32类型仅能保留约6-7位有效数字,部分十进制小数无法被二进制浮点格式精确表示,舍入操作会触发微小的截断误差。
- 底层实现差异: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
相关产品推荐
相关产品推荐

