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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 04:25:18