ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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

JAX 类型提升怎么判断结果 dtype:弱类型与 promote_types 类型格 JAX 类型提升怎么判断结果 dtype弱类型与 promote_types 类型格【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax在 JAX 里写数值计算时一个常见任务是先确认二元运算的结果 dtype例如jnp.float32数组加一个 Pythonint到底得到什么类型混合了int8和float16的表达式会提升到哪种精度JAX 的类型提升规则与 NumPy 不完全一致直接照搬 NumPy 的直觉容易判断错。用 JAX 判断结果 dtype 的路径有两条一是对类型对象调用jnp.promote_types直接查询结果类型二是查 JAX 的类型提升格type promotion lattice图结合弱类型weak type规则人工推演。本文介绍这两条路径怎么用、JAX 与 NumPy 规则差异在哪里以及如何用 strict 提升模式在代码里强制校验自己的类型判断。类型提升格如何读出结果 dtypeJAX 的类型提升规则由一张类型提升格定义官方说明见 docs/101/type_promotion.rst格图本身是 docs/_static/type_lattice.svg。任意两个类型组合后的结果类型就是这两个类型在这张格上的 join最小上界。图中节点用短名表示 dtype对应关系如下均来自 docs/101/type_promotion.rstb1表示np.bool_i2表示np.int16u4表示np.uint32bf表示np.bfloat16f2表示np.float16c8表示np.complex64i*表示 Pythonint或弱类型的intf*表示 Pythonfloat或弱类型的floatc*表示 Pythoncomplex或弱类型的complex。推演示例要判断uint8和int16运算后的结果在格上找到u1和i1它们的最小上界是i2即结果是int16。带星号的i*、f*、c*是弱类型节点下一节会解释为什么它们几乎总是“被对面吃掉”。用 jnp.promote_types 直接查询不想查格图时可以直接调用jnp.promote_types(a, b)。它是numpy.promote_types的 JAX 实现返回二元运算应将参数转换到的类型源码见 jax/_src/dtypes.py。参数a和b可以传字符串、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)内置标量类型int、float、complex被当作弱类型处理不会改变强类型对方的位宽而 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 用带绿色背景的单元格标出了所有差异位置归纳起来是三类弱类型对强类型同类别时JAX 总是优先保留 JAX 值的精度。例如jnp.int16(1) 1返回int16而不是 NumPy 会提升到的int64。注意这条只适用于 Python 标量如果常量是 NumPy 数组则按格走jnp.int16(1) np.array(1)返回int64。整数/布尔对浮点或复数时JAX 总是优先浮点/复数一方。例如int8与float16运算得到float16而不是像经典 NumPy 规则那样推向 64 位浮点。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, dtypeint8) 2 * x Array([0, 2, 4, 6, 8], dtypeint8)弱类型框架的目的就是防止 JAX 值与“没有显式指定类型的值”如 Python 标量字面量做二元运算时发生不想要的提升。如果2不被当作弱类型上面的表达式就会被提升 jnp.int32(2) * x Array([0, 2, 4, 6, 8], dtypeint32)Python 标量在 JAX 中有时会被提升为 DeviceArray 对象例如 JIT 编译期间。为了在这种情况下仍保持提升语义DeviceArray 带有一个weak_type标志可以直接从数组的字符串表示中看出来 jnp.asarray(2) Array(2, dtypeint32, weak_typeTrue)显式指定dtype则得到强类型数组 jnp.asarray(2, dtypeint32) Array(2, dtypeint32)所以核对一个值的类型身份时除了看dtype还要看weak_type标志Array(2, dtypeint32, weak_typeTrue)和Array(2, dtypeint32)参与后续运算的行为是不同的——前者按弱类型走格上的i*节点后者按i4节点走。用 strict 提升模式校验类型假设如果对隐式提升不放心可以把隐式提升关掉要求所有提升显式进行把jax_numpy_dtype_promotion设为strict。该配置只有两个取值standard默认和strictstrict 模式下两个强指定 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_promotionstrict. Try explicitly casting inputs to the desired output type, or set jax_numpy_dtype_promotionstandard.注意 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 的完整闭环可以这样落地对确定的 dtype 对用jnp.promote_types(a, b)查询返回的dtype(...)就是预期结果类型涉及 Python 标量时先确认该值是弱类型看weak_typeTrue标志或按i*/f*/c*节点处理它不会抬高对面 JAX 值的精度实际运行后核对结果的dtypeArray(..., dtype...)字符串表示中可见与第 1 步的查询结果一致即说明判断成立需要防止隐式提升悄悄改变 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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表