默认数据类型与 X64 标志#
JAX 致力于满足不同数值计算从业者的需求,而这些需求有时会存在冲突。在默认数据类型(dtype)方面,存在两个不同的阵营:
经典的科学计算从业者(即
numpy或scipy等工具的用户)往往将计算的精度放在首位:这类用户倾向于让计算默认使用可用的最宽表示:例如,浮点值应默认为float64,整数应默认为int64等。人工智能研究人员(即实现和训练神经网络的人员)往往更看重速度而非精度,甚至开发了如 bfloat16 等特殊数据类型,通过刻意舍弃最低有效位来加速计算。对于这些用户而言,计算中出现 float64 值往好了说是程序变慢,往坏了说甚至会导致与硬件不兼容!这些用户更希望计算默认使用
float32或int32。
JAX 为此提供的主要机制是 jax_enable_x64 标志,它控制是否允许创建 64 位值。该标志默认设置为 False(以满足 AI 研究人员和从业者的需求),但对于重视精度胜过计算速度的用户,可以将其设置为 True。
默认设置:全局 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)
除了默认值之外,由于 64 位值对于 AI 工作流来说可能具有“毒性”,将此标志设为 False 还可以防止创建任何 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. See https://github.com/jax-ml/jax#current-gotchas for more.
Array([0., 1., 2., 3., 4.], dtype=float32)
X64 标志:启用 64 位值#
若要在“另一种模式”下工作(即函数默认产生 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)
X64 配置也可以通过 shell 环境变量 JAX_ENABLE_X64 进行设置,例如:
$ JAX_ENABLE_X64=1 python main.py
X64 标志旨在作为一种全局设置,在整个程序中应保持统一的值,并应在主文件的顶部进行设置。一个常见的需求是希望该标志可以按上下文配置(例如仅在长程序的某一部分启用 X64):事实证明,这在 JAX 的编程模型中很难实现,因为代码的执行可能发生在与代码编译不同的上下文中。目前我们正在探索放宽此约束的可行性,敬请期待!