JAX浮点数精度问题:输入微变时jnp.arctan结果无变化如何解决?
JAX中arctan对微小输入变化无响应的解决方法
问题本质
你遇到的情况是JAX默认使用单精度浮点数(float32),而Numpy默认用**双精度(float64)**导致的:
- float32的有效精度只有6-7位,10和10+1e-7这两个值在float32的精度范围内完全相等——1e-7相对于10的比例太小,超出了float32的分辨能力,所以JAX的
arctan输出自然没有变化。 - float64的有效精度能到15-16位,足以区分这两个微小差异,因此Numpy的结果会有区别。
解决方法
1. 全局开启JAX的64位精度模式
在代码开头添加配置,让JAX默认使用float64计算:
import jax jax.config.update("jax_enable_x64", True) import numpy as np import jax.numpy as jnp # 再运行测试代码就能得到和Numpy一致的结果 print('jnp.arctan(10) is:','%.60f' % jnp.arctan(10)) print('np.arctan(10) is:','%.60f' % np.arctan(10))
2. 针对单个操作显式指定float64类型
如果不想全局修改精度,可以在输入时强制使用双精度数组:
import numpy as np import jax.numpy as jnp # 显式将输入转为float64 print('jnp.arctan(10+1e-7) is:','%.60f' % jnp.arctan(jnp.array(10+1e-7, dtype=jnp.float64))) print('np.arctan(10+1e-7) is:','%.60f' % np.arctan(10+1e-7))
3. 验证输入的精度差异
你可以先确认输入在不同精度下是否真的相等:
import jax.numpy as jnp x_float32_1 = jnp.array(10, dtype=jnp.float32) x_float32_2 = jnp.array(10+1e-7, dtype=jnp.float32) print(x_float32_1 == x_float32_2) # 输出True,说明float32下两者无差异 x_float64_1 = jnp.array(10, dtype=jnp.float64) x_float64_2 = jnp.array(10+1e-7, dtype=jnp.float64) print(x_float64_1 == x_float64_2) # 输出False,float64能区分差异
内容的提问来源于stack exchange,提问作者funpy
相关产品推荐
相关产品推荐

