首页
/ JAX 默认 dtype 与 X64 标志:理解 jax_enable_x64 的工作原理与使用方式

JAX 默认 dtype 与 X64 标志:理解 jax_enable_x64 的工作原理与使用方式

2026-09-05 19:59:52作者:乔或婵

JAX 需要在"以精度为先"的科学计算用户和"以速度为先"的 AI 训练用户之间做出权衡,这一权衡集中体现在 jax_enable_x64 配置标志上。本文基于 JAX 官方文档 Default dtypes and the X64 flag 展开,结合 配置定义dtype 规范化实现 的源码,讲清楚 JAX 默认 32 位行为的由来、显式请求 64 位类型时的处理路径,以及启用 X64 模式的所有官方方式与限制条件。

两类用户,两种默认偏好

JAX 的用户群体对默认 dtype 的期望存在明显分歧,文档将其归纳为两个阵营:

  • 经典科学计算用户(numpy / scipy 类工具的使用者)把计算精度放在第一位,期望运算默认使用当前平台上最宽的表示:浮点数默认 float64、整数默认 int64;
  • AI 研究者(实现和训练神经网络的人)把速度放在精度之上,甚至专门发展出 bfloat16 这类刻意丢弃低位有效比特以换取计算速度的数据类型。对这类用户而言,计算中一旦出现 float64 值,轻则导致程序变慢,重则与硬件不兼容。他们期望默认 dtype 是 float32int32

JAX 提供的核心机制就是 jax_enable_x64 标志:它控制程序能否创建 64 位数值。默认值为 False(服务 AI 研究者),重视精度的用户可将其设为 True

从源码看,该标志在 jax/_src/config.py 中定义:

enable_x64 = bool_state(
    name='jax_enable_x64',
    default=False,
    help='Enable 64-bit types to be used',
    include_in_jit_key=True,
    include_in_trace_context=True)

两个细节值得注意:一是 include_in_jit_key=True,意味着 X64 模式会参与 JIT 编译缓存键的计算,切换该标志会使已编译产物失效并按新位宽重新编译;二是同一文件中还注册了一个便捷只读属性 jax/_src/config.py:

setattr(Config, "x64_enabled", property(lambda _: enable_x64.value))

因此运行期可以用 jax.config.x64_enabled 查询当前是否启用了 64 位模式。

默认设置:处处 32 位

jax_enable_x64 默认为 False,因此 jax.numpy 的数组创建函数默认返回 32 位值:

>>> import jax.numpy as jnp

>>> jnp.arange(5)
Array([0, 1, 2, 3, 4], dtype=int32)

>>> jnp.zeros(5)
Array([0., 0., 0., 0., 0.], dtype=float32)

>>> jnp.ones(5, dtype=int)
Array([1, 1, 1, 1, 1], dtype=int32)

最后一个例子最能体现 JAX 与 NumPy 的差异:在 Python 中 int 对应 64 位,但 JAX 默认模式下 jnp.ones(5, dtype=int) 得到的是 int32

这背后的机制在 jax/_src/dtypes.py 中可以完整看到。首先,Python 标量类型到默认 dtype 的映射表 jax/_src/dtypes.py 与 NumPy 相同(intint64,floatfloat64):

python_scalar_types_to_dtypes: dict[type, DType] = {
  bool: np.dtype('bool'),
  int: np.dtype('int64'),
  float: np.dtype('float64'),
  complex: np.dtype('complex128'),
}

随后统一经过 canonicalize_dtype 规范化,而规范化函数内部维护了一张 64 位 → 32 位的降级表 jax/_src/dtypes.py:

_dtype_to_32bit_dtype: dict[DType, DType] = {
    np.dtype('int64'): np.dtype('int32'),
    np.dtype('uint64'): np.dtype('uint32'),
    np.dtype('float64'): np.dtype('float32'),
    np.dtype('complex128'): np.dtype('complex64'),
}

核心逻辑 jax/_src/dtypes.py:

@functools.cache
def _canonicalize_dtype(x64_enabled: bool, allow_extended_dtype: bool, dtype: Any) -> DType | ExtendedDType:
  ...
  if x64_enabled:
    return dtype_
  else:
    return _dtype_to_32bit_dtype.get(dtype_, dtype_)

也就是说,jax_enable_x64 关闭时,任何映射表中出现的 64 位 dtype 都会被静默降级为对应 32 位类型;表外的类型(如本来就是 32 位或更窄的类型)原样保留。

显式请求 64 位:默认行为是警告并截断

文档特别强调:由于 64 位值对 AI 工作流"毒性太大",jax_enable_x64=False 时 JAX 阻止你创建 64 位数组。显式请求时的示例输出为:

>>> jnp.arange(5, dtype='float64')
UserWarning: Explicitly requested dtype float64 requested in arange is not available, and will be
truncated to dtype float32. To enable more dtypes, set the jax_enable_x64 configuration option or the
JAX_ENABLE_X64 shell environment variable.
Array([0., 1., 2., 3., 4.], dtype=float32)

这段警告并非凭空产生,它来自 jax/_src/dtypes.py 中的 _maybe_canonicalize_explicit_dtype 函数,该函数还会读取另一个配套配置 jax_explicit_x64_dtypes,决定显式 64 位请求的处理策略。该配置定义在 jax/_src/config.py,有三种取值:

取值 行为
allow 即使 enable_x64 为 False,也尊重显式指定的 64 位类型
warn(默认) 发出警告,并将类型截断为对应 32 位类型
error 直接抛出 ValueError

源码分支 jax/_src/dtypes.py:

allow = config.explicit_x64_dtypes.value
if allow == config.ExplicitX64Mode.ALLOW or config.enable_x64.value:
    return dtype
canonical_dtype = canonicalize_dtype(dtype)
...
if allow == config.ExplicitX64Mode.ERROR:
    msg = ("Explicitly requested dtype {}{} is not available. ...")
    raise ValueError(msg)
else:  # WARN
    msg = ("Explicitly requested dtype {}{} is not available, "
          "and will be truncated to dtype {}. ...")
    warnings.warn(msg, stacklevel=4)
    return canonical_dtype

这对集成第三方库的用户尤其有用:若外部代码偶尔传入 float64 参数,默认的 warn 模式会安全降级而不崩溃;而追求严格性的项目可以把 jax_explicit_x64_dtypes 设为 error,让任何"漏网"的 64 位请求在开发阶段就暴露出来。

启用 64 位:X64 标志

要切换到"函数默认产出 64 位值"的模式,把 jax_enable_x64 设为 True 即可:

import jax
import jax.numpy as jnp

jax.config.update('jax_enable_x64', True)

print(repr(jnp.arange(5)))
print(repr(jnp.zeros(5)))
print(repr(jnp.ones(5, dtype=int)))
Array([0, 1, 2, 3, 4], dtype=int64)
Array([0., 0., 0., 0., 0.], dtype=float64)
Array([1, 1, 1, 1, 1], dtype=int64)

对照前文的 canonicalize_dtype 源码即可理解其效果:开启后 _canonicalize_dtype 直接返回请求的 dtype,不再查降级表,于是 jnp.arange 的整数默认值从 int32 变为 int64,jnp.zeros 的浮点默认值从 float32 变为 float64

X64 配置也可以通过 shell 环境变量 JAX_ENABLE_X64 设置,例如:

$ JAX_ENABLE_X64=1 python main.py

对于启动脚本、容器镜像或 CI 环境,环境变量方式比在代码里 jax.config.update 更可靠,因为它在 Python 进程任何代码执行之前就已生效。

X64 标志是全局设置:为什么不能"只开一段"

文档明确指出:X64 标志被设计为全局设置,应当对整个程序保持同一个值,并在主文件顶部一次性设置。一个常见的功能请求是让它支持上下文级配置(例如只在长程序的某一段开启 X64),但这在 JAX 的编程模型中难以实现——因为代码的执行可能与编译发生在不同的上下文中。文档表示有正在进行的工作在探索放宽这一限制的可行性。

源码中也印证了"全局"这一设计立场。jax/_src/config.py 在创建 enable_x64 状态之后,紧接着把它从"可用上下文管理器控制的标志"列表中移除:

jax_jit.set_enable_x64_state(enable_x64)

# TODO(phawkins): remove after fixing users of FLAGS.x64_enabled.
config._contextmanager_flags.remove('jax_enable_x64')

这意味着 jax_enable_x64 不在 JAX 支持上下文管理器切换的配置项之列,试图用 with jax.config.use(...) 的方式局部改写它是不被支持的;正确做法是在程序入口统一设置。

另外,由于该标志参与了 JIT 缓存键(include_in_jit_key=True),同一份代码在 X64 开/关两种模式下会分别编译、分别缓存,互不污染。相关的行为验证可以参见 x64 上下文测试,它专门覆盖不同 x64 模式下的 dtype 行为。

实践小结

  • AI 训练/推理代码:保持默认 jax_enable_x64=False,全程 float32/int32;如需更高吞吐,可另行使用 bfloat16 等低精度类型,而不是回到 float64。
  • 科学计算代码:在 main.py 顶部执行 jax.config.update('jax_enable_x64', True),或部署时统一设置 JAX_ENABLE_X64=1;需要精确控制显式 64 位请求时,配合 jax_explicit_x64_dtypeswarn(默认)/error/allow 三种模式。
  • 排查意外 dtype 时:记住"映射表降级"这一条路径——python_scalar_types_to_dtypes 先按 NumPy 习惯映射到 64 位,再由 canonicalize_dtype 依据 enable_x64 决定是否截断,这是 JAX 与 NumPy 默认 dtype 差异的根源,也对应文档示例中 dtype=int 得到 int32 的现象。

以上机制与行为均以当前仓库 jax/_src/config.pyjax/_src/dtypes.py 的源码为准;若在 X64 关闭模式下看到"Explicitly requested dtype ... is not available, and will be truncated"警告,按警告提示开启 X64 或调整 jax_explicit_x64_dtypes 即可。

登录后查看全文
热门项目推荐
相关项目推荐