9.536743e-7的含义是什么?为何会出现在jax.numpy代码中?
关于9.536743e-7数值的意义与JAX中离散化问题的解答
数值实际意义
- 你观察到的9.536743e-7本质是2的负20次方的近似值,精确值为
1/(2^20) = 1/1048576 ≈ 9.5367431640625e-7,和你搜索到的所谓比特/兆比特换算完全无关,相关说法是混淆二进制存储单位和浮点数精度的错误内容。 - 该数值是单精度浮点数(float32)在数值1附近的最小可表示间隔,也就是单精度浮点数的ULP(Unit in the Last Place,最低有效位对应单位值):float32的尾数部分占23个二进制位,加上隐含的1位整数位,共24位有效精度,对于量级为1的数值,相邻两个可表示的float32数值的差恰好等于2^-20。
该数值出现在JAX代码中的原因
- JAX默认关闭float64支持,所有浮点运算默认采用float32精度,和NumPy默认使用float64的行为不同。
- 当你的函数输出值量级在1附近时,float32最小只能表示约9.5e-7的变化,小于这个步长的数值变化会被浮点舍入规则抹除,最终表现为输出呈现离散化的特征。
验证与解决方法
- 你可以通过以下代码验证该现象:
import jax.numpy as jnp # 输出为False,说明加该数值刚好可以让1的float32表示发生变化 print(jnp.float32(1) + 9.536743e-7 == jnp.float32(1)) # 输出为True,说明加小于该数值的量无法改变1的float32表示 print(jnp.float32(1) + 5e-7 == jnp.float32(1))
- 如果需要消除离散化问题,开启JAX的float64支持即可,在代码入口处添加配置:
import jax jax.config.update("jax_enable_x64", True)
开启后浮点运算默认采用float64精度,对应1附近的最小可表示间隔缩小到约2.2e-16,常规场景下不会再观察到离散化现象。
内容的提问来源于stack exchange,提问作者AetbeUT
相关产品推荐
相关产品推荐

