【问题标题】:Why does Mypy think adding two Jax arrays returns a numpy array?为什么 Mypy 认为添加两个 Jax 数组会返回一个 numpy 数组?
【发布时间】:2021-08-22 18:44:40
【问题描述】:

考虑以下文件:

import jax.numpy as jnp

def test(a: jnp.ndarray, b: jnp.ndarray) -> jnp.ndarray:
    return a + b

运行mypy mypytest.py 返回以下错误:

mypytest.py:4: error: Incompatible return value type (got "numpy.ndarray[Any, dtype[bool_]]", expected "jax._src.numpy.lax_numpy.ndarray")

由于某种原因,它认为添加两个 jax.numpy.ndarrays 会返回一个 bools 的 NumPy 数组。难道我做错了什么?或者这是 MyPy 中的错误,还是 Jax 的类型注释?

【问题讨论】:

    标签: python type-hinting mypy python-typing jax


    【解决方案1】:

    至少在静态上,jnp.ndarraynp.ndarray 的子类,修改非常少

    class ndarray(np.ndarray, metaclass=_ArrayMeta):
      dtype: np.dtype
      shape: Tuple[int, ...]
      size: int
    
      def __init__(shape, dtype=None, buffer=None, offset=0, strides=None,
                   order=None):
        raise TypeError("jax.numpy.ndarray() should not be instantiated explicitly."
                        " Use jax.numpy.array, or jax.numpy.zeros instead.")
    

    因此,它继承了np.ndarray 的方法类型签名。

    我猜运行时行为是通过jnp.array 函数实现的。除非我错过了一些存根文件或类型诡计,否则jnp.array 的结果与jnp.ndarray 匹配只是因为jnp.array 没有类型化。您可以使用

    进行测试
    def foo(_: str) -> None:
       pass
    
    foo(jnp.array(0))
    

    通过 mypy。

    所以回答你的问题,我不认为你做错了什么。这是一个错误,它可能不是他们的意思,但它实际上并不是不正确的,因为当您添加 jnp.ndarrays 时,您确实得到了 np.ndarray,因为 jnp.ndarraynp.ndarray

    至于为什么bools,那很可能是因为你的jnp.arrays 缺少泛型参数,而np.ndarray__add__ 的第一个有效重载是

        @overload
        def __add__(self: NDArray[bool_], other: _ArrayLikeBool_co) -> NDArray[bool_]: ...  # type: ignore[misc]
    

    所以它只是默认为bool

    【讨论】:

      【解决方案2】:

      一般来说,JAX 与 mypy 的兼容性很差,因为 JAX 的转换模型很难满足 mypy 的约束,它经常调用具有转换特定跟踪器值的函数,这些跟踪器值充当数组的替身(参见 How To Think in JAX: JIT Mechanics简要讨论此机制)。

      使用跟踪器类型作为数组的替代意味着 mypy 在转换严格类型的 JAX 函数时会引发错误,因此在整个 JAX 代码库中,我们倾向于将 Array 别名为 Any,并使用此作为返回数组的 JAX 函数的返回类型注释。

      在这方面改进会很好,因为Any 返回类型对于有效的类型检查不是很有用,但这只是使 mypy 与 JAX 良好配合的众多挑战中的第一个。如果你想阅读过去几年围绕这个问题的一些讨论,我会从这里开始:https://github.com/google/jax/issues/943

      同时,我的建议是使用 Any 作为 JAX 数组的类型注释。

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 1970-01-01
        • 2016-09-16
        • 1970-01-01
        • 2018-04-22
        • 1970-01-01
        • 1970-01-01
        • 2020-07-07
        相关资源
        最近更新 更多