JAX 类型提升怎么判断结果 dtype:弱类型与 promote_types 类型格
2026/9/11 17:52:40 网站建设 项目流程

JAX 类型提升怎么判断结果 dtype:弱类型与 promote_types 类型格

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

在 JAX 里写数值计算时,一个常见任务是先确认二元运算的结果 dtype:例如jnp.float32数组加一个 Pythonint到底得到什么类型?混合了int8float16的表达式会提升到哪种精度?JAX 的类型提升规则与 NumPy 不完全一致,直接照搬 NumPy 的直觉容易判断错。用 JAX 判断结果 dtype 的路径有两条:一是对类型对象调用jnp.promote_types直接查询结果类型;二是查 JAX 的类型提升格(type promotion lattice)图,结合弱类型(weak type)规则人工推演。本文介绍这两条路径怎么用、JAX 与 NumPy 规则差异在哪里,以及如何用 strict 提升模式在代码里强制校验自己的类型判断。

类型提升格:如何读出结果 dtype

JAX 的类型提升规则由一张类型提升格定义,官方说明见 docs/101/type_promotion.rst,格图本身是 docs/_static/type_lattice.svg。

任意两个类型组合后的结果类型,就是这两个类型在这张格上的 join(最小上界)。图中节点用短名表示 dtype,对应关系如下(均来自 docs/101/type_promotion.rst):

  • b1表示np.bool_
  • i2表示np.int16
  • u4表示np.uint32
  • bf表示np.bfloat16
  • f2表示np.float16
  • c8表示np.complex64
  • i*表示 Pythonint或弱类型的int
  • f*表示 Pythonfloat或弱类型的float
  • c*表示 Pythoncomplex或弱类型的complex

推演示例:要判断uint8int16运算后的结果,在格上找到u1i1,它们的最小上界是i2,即结果是int16。带星号的i*f*c*是弱类型节点,下一节会解释为什么它们几乎总是“被对面吃掉”。

用 jnp.promote_types 直接查询

不想查格图时,可以直接调用jnp.promote_types(a, b)。它是numpy.promote_types的 JAX 实现,返回二元运算应将参数转换到的类型(源码见 jax/_src/dtypes.py)。

参数ab可以传字符串、dtype 对象或标量类型,返回值始终是numpy.dtype。文档中给出的示例:

>>> import jax.numpy as jnp >>> jnp.promote_types('int32', 'float32') # strings dtype('float32') >>> jnp.promote_types(jnp.dtype('int32'), jnp.dtype('float32')) # dtypes dtype('float32') >>> jnp.promote_types(jnp.int32, jnp.float32) # scalar types dtype('float32')

内置标量类型(intfloatcomplex)被当作弱类型处理,不会改变强类型对方的位宽;而 NumPy 的同名函数把这些类型视作 64 位类型,这是两者的关键差异:

>>> jnp.promote_types('uint8', int) dtype('uint8') >>> jnp.promote_types('float16', float) dtype('float16') >>> import numpy >>> numpy.promote_types('uint8', int) dtype('int64') >>> numpy.promote_types('float16', float) dtype('float64')

这段输出的来源是jnp.promote_types文档字符串中的示例(见 jax/_src/dtypes.py)。也就是说:在 JAX 里,Python 标量作为“弱类型”参与运算时保持对面 JAX 值的精度,不会像 NumPy 那样一律推到 64 位。

JAX 与 NumPy 提升规则的三类差异

jnp.promote_types的结果与numpy.promote_types不完全相同。docs/101/type_promotion.rst 用带绿色背景的单元格标出了所有差异位置,归纳起来是三类:

  1. 弱类型对强类型(同类别)时,JAX 总是优先保留 JAX 值的精度。例如jnp.int16(1) + 1返回int16而不是 NumPy 会提升到的int64。注意这条只适用于 Python 标量;如果常量是 NumPy 数组,则按格走:jnp.int16(1) + np.array(1)返回int64
  2. 整数/布尔对浮点或复数时,JAX 总是优先浮点/复数一方。例如int8float16运算得到float16,而不是像经典 NumPy 规则那样推向 64 位浮点。
  3. bfloat16 的行为。JAX 支持非标准 16 位浮点类型jax.numpy.bfloat16,对神经网络训练有用。它唯一值得注意的提升行为是与 IEEE-754float16:两者混合时提升到float32

文档给出的动机是:GPU 使用 64 位浮点代价很高,TPU 干脆不支持 64 位浮点,经典 NumPy 规则“太愿意提升到 64 位”,不适合面向加速器的系统。JAX 的浮点提升规则更保守,与 PyTorch 的规则相似。判断结果 dtype 时,凡是会“意外变大到 64 位”的地方,基本都能在这三类差异里找到原因。

注意 Python 运算符的派发改写规则

还有一个容易踩的边界:Python 运算符(如+)按操作数的 Python 类型来派发规则。因此np.int16(1) + 1按 NumPy 规则提升,而jnp.int16(1) + 1按 JAX 规则提升。一旦两种规则混在一个表达式里,就可能出现不符合直觉的非结合性提升语义,例如np.int16(1) + 1 + jnp.int16(1)。判断 dtype 前先确认每个操作数到底是 JAX 值、NumPy 值还是 Python 标量,再决定用哪套规则。

弱类型:识别并核对 weak_type 标志

JAX 的弱类型(weak type)值在大多数情况下可以当作 Python 标量对待。文档示例:

>>> import jax.numpy as jnp >>> x = jnp.arange(5, dtype='int8') >>> 2 * x Array([0, 2, 4, 6, 8], dtype=int8)

弱类型框架的目的就是防止 JAX 值与“没有显式指定类型的值”(如 Python 标量字面量)做二元运算时发生不想要的提升。如果2不被当作弱类型,上面的表达式就会被提升:

>>> jnp.int32(2) * x Array([0, 2, 4, 6, 8], dtype=int32)

Python 标量在 JAX 中有时会被提升为 DeviceArray 对象(例如 JIT 编译期间)。为了在这种情况下仍保持提升语义,DeviceArray 带有一个weak_type标志,可以直接从数组的字符串表示中看出来:

>>> jnp.asarray(2) Array(2, dtype=int32, weak_type=True)

显式指定dtype则得到强类型数组:

>>> jnp.asarray(2, dtype='int32') Array(2, dtype=int32)

所以核对一个值的类型身份时,除了看dtype,还要看weak_type标志:Array(2, dtype=int32, weak_type=True)Array(2, dtype=int32)参与后续运算的行为是不同的——前者按弱类型走格上的i*节点,后者按i4节点走。

用 strict 提升模式校验类型假设

如果对隐式提升不放心,可以把隐式提升关掉,要求所有提升显式进行:把jax_numpy_dtype_promotion设为'strict'。该配置只有两个取值,standard(默认)和strict;strict 模式下,两个强指定 dtype 不同的数组做二元运算会直接报错。配置定义见 jax/_src/config.py。

局部启用用上下文管理器jax.numpy_dtype_promotion('strict')。文档示例(含文档中给出的报错输出):

>>> import jax >>> import jax.numpy as jnp >>> x = jnp.float32(1) >>> y = jnp.int32(1) >>> with jax.numpy_dtype_promotion('strict'): ... z = x + y Traceback (most recent call last): TypePromotionError: Input dtypes ('float32', 'int32') have no available implicit dtype promotion path when jax_numpy_dtype_promotion=strict. Try explicitly casting inputs to the desired output type, or set jax_numpy_dtype_promotion=standard.

注意 strict 模式仍然允许“安全的弱类型提升”,JAX 数组与 Python 标量混合的代码照常可写:

>>> with jax.numpy_dtype_promotion('strict'): ... z = x + 1 >>> print(z) 2.0

想全局启用就用标准配置更新接口,恢复默认同理:

jax.config.update('jax_numpy_dtype_promotion', 'strict') # 恢复默认的 standard 提升 jax.config.update('jax_numpy_dtype_promotion', 'standard')

验证方式与边界

判断结果 dtype 的完整闭环可以这样落地:

  1. 对确定的 dtype 对,用jnp.promote_types(a, b)查询,返回的dtype(...)就是预期结果类型;
  2. 涉及 Python 标量时,先确认该值是弱类型(看weak_type=True标志或按i*/f*/c*节点处理),它不会抬高对面 JAX 值的精度;
  3. 实际运行后核对结果的dtypeArray(..., dtype=...)字符串表示中可见),与第 1 步的查询结果一致即说明判断成立;
  4. 需要防止隐式提升悄悄改变 dtype 的代码,用jax.numpy_dtype_promotion('strict')上下文或jax.config.update('jax_numpy_dtype_promotion', 'strict')把隐式提升变成TypePromotionError,报错信息会指出是哪两个 dtype(如('float32', 'int32'))没有隐式提升路径,并提示显式 cast 或切回standard

边界情况有两个,直接来自 docs/101/type_promotion.rst:一是弱类型规则只针对 Python 标量,NumPy 数组常量仍按格提升(jnp.int16(1) + np.array(1)int64);二是 NumPy 值与 JAX 值混用同一表达式时,运算符派发会交替套用两套规则,产生非结合性的提升行为。相关设计背景可继续查阅仓库内的 docs/jep/9407-type-promotion.md。

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询