diff --git a/etils/enp/array_spec.py b/etils/enp/array_spec.py index 3b7aeaea..23dda57a 100644 --- a/etils/enp/array_spec.py +++ b/etils/enp/array_spec.py @@ -165,7 +165,7 @@ def is_fake_array(array: Array) -> bool: # pyrefly: ignore[not-a-type] def _is_flax_summary(value: Array) -> bool: # pyrefly: ignore[not-a-type] if 'flax.linen' not in sys.modules: return False - from flax import linen as nn # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from flax import linen as nn # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return isinstance(value, nn.summary._ArrayRepresentation) # pylint: disable=protected-access @@ -173,15 +173,15 @@ def _is_flax_summary(value: Array) -> bool: # pyrefly: ignore[not-a-type] def _is_grain(array: Array) -> bool: # pyrefly: ignore[not-a-type] if 'grain.tensorflow' not in sys.modules: return False - from grain import tensorflow as grain # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from grain import tensorflow as grain # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return isinstance(array, grain.ArraySpec) def _get_grain_shm_array_metadata_cls(): """Imports the shm metadata from `grain` in a cross-version compatible way.""" - # pylint: disable=g-import-not-at-top # pytype: disable=import-error - import grain + # pylint: disable=g-import-not-at-top + import grain # pyrefly: ignore[missing-import] cls = getattr( grain.multiprocessing, @@ -189,10 +189,10 @@ def _get_grain_shm_array_metadata_cls(): None, ) if cls is None: - from grain._src.python import shared_memory_array + from grain._src.python import shared_memory_array # pyrefly: ignore[missing-import] cls = shared_memory_array.SharedMemoryArrayMetadata - # pylint: enable=g-import-not-at-top # pytype: enable=import-error + # pylint: enable=g-import-not-at-top return cls @@ -209,7 +209,7 @@ def _is_pygrain(array: Array) -> bool: # pyrefly: ignore[not-a-type] def _is_orbax(array: Array) -> bool: # pyrefly: ignore[not-a-type] if 'orbax.checkpoint' not in sys.modules: return False - from orbax.checkpoint.metadata import value # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from orbax.checkpoint.metadata import value # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return isinstance( array, diff --git a/etils/enp/checking.py b/etils/enp/checking.py index 889c5005..ce8a3dc8 100644 --- a/etils/enp/checking.py +++ b/etils/enp/checking.py @@ -248,7 +248,7 @@ def _maybe_set_tnp_casting(xnp: numpy_utils.NpModule) -> None: # When TF numpy mode is not enabled, `tnp.asarray(1.)` returns tf.float64, # creating conflict because TF do fail for operations like: # `tf.float64 + tf.float32` - from tensorflow.python.ops.numpy_ops import np_dtypes # pylint: disable=g-import-not-at-top,g-direct-tensorflow-import # pytype: disable=import-error + from tensorflow.python.ops.numpy_ops import np_dtypes # pylint: disable=g-import-not-at-top,g-direct-tensorflow-import # pyrefly: ignore[missing-import] if not np_dtypes.is_prefer_float32(): np_dtypes.set_prefer_float32(True) diff --git a/etils/enp/checking_test.py b/etils/enp/checking_test.py index 7740bfe8..8c15309c 100644 --- a/etils/enp/checking_test.py +++ b/etils/enp/checking_test.py @@ -79,11 +79,11 @@ def test_type(xnp: enp.NpModule, fn): fn(x, xnp.asarray(2.0, dtype=xnp.float32)) # Independently of the original xnp, we can explicitly pass the target xnp - _assert_out(fn(x, y, xnp=enp.lazy.np), enp.lazy.np) # pytype: disable=wrong-keyword-args - _assert_out(fn(x, y, xnp=enp.lazy.jnp), enp.lazy.jnp) # pytype: disable=wrong-keyword-args - _assert_out(fn(x, y, xnp=enp.lazy.tnp), enp.lazy.tnp) # pytype: disable=wrong-keyword-args + _assert_out(fn(x, y, xnp=enp.lazy.np), enp.lazy.np) + _assert_out(fn(x, y, xnp=enp.lazy.jnp), enp.lazy.jnp) + _assert_out(fn(x, y, xnp=enp.lazy.tnp), enp.lazy.tnp) # TODO(epot): `torch.asarray` do not work with `tf` / `jax` - # _assert_out(fn(x, y, xnp=enp.lazy.torch), enp.lazy.torch) # pytype: disable=wrong-keyword-args + # _assert_out(fn(x, y, xnp=enp.lazy.torch), enp.lazy.torch) # Pass a xnp and np yield xnp _assert_out(fn(x, np.asarray(y)), xnp) @@ -95,21 +95,21 @@ def test_type(xnp: enp.NpModule, fn): def test_non_array_annotations(): @enp.check_and_normalize_arrays(strict=False) - def fn_non_array_args(x: int, y: FloatArray, z): + def fn_non_array_args(x: int, y: FloatArray, z): # pyrefly: ignore[not-a-type] # Non-array typing annotations are preserved assert isinstance(x, int) assert isinstance(z, str) assert enp.lazy.get_xnp(y) is enp.lazy.jnp return y + x - _assert_out(fn_non_array_args(1, [2], 'abc', xnp=enp.lazy.jnp), enp.lazy.jnp) # pytype: disable=wrong-keyword-args + _assert_out(fn_non_array_args(1, [2], 'abc', xnp=enp.lazy.jnp), enp.lazy.jnp) # pyrefly: ignore[unexpected-keyword] def test_missing_xnp_default(): @enp.check_and_normalize_arrays(strict=False) - def fn_missing_default(x: FloatArray, *, xnp: enp.NpModule): + def fn_missing_default(x: FloatArray, *, xnp: enp.NpModule): # pyrefly: ignore[not-a-type] del xnp return x - fn_missing_default(1.0) # pytype: disable=missing-parameter # pylint: disable=missing-kwoa + fn_missing_default(1.0) # pylint: disable=missing-kwoa # pyrefly: ignore[missing-argument] fn_missing_default(1.0, xnp=np) diff --git a/etils/enp/compat.py b/etils/enp/compat.py index cf15885e..7cbcb423 100644 --- a/etils/enp/compat.py +++ b/etils/enp/compat.py @@ -32,7 +32,7 @@ import numpy as np if typing.TYPE_CHECKING: - import torch as torch_ # pytype: disable=import-error + import torch as torch_ # pyrefly: ignore[missing-import] _NpDType = Any diff --git a/etils/enp/interp_utils.py b/etils/enp/interp_utils.py index 25a0ca97..9cc70cef 100644 --- a/etils/enp/interp_utils.py +++ b/etils/enp/interp_utils.py @@ -106,7 +106,7 @@ def interp( to = tuple(xnp.asarray(v) for v in to) # `a` can be scalar or array of shape=(x.shape[-1],), same for `b` - a, b = _linear_interp_factors(*from_, *to) # pytype: disable=wrong-arg-types + a, b = _linear_interp_factors(*from_, *to) return a * x + b diff --git a/etils/enp/numpy_utils.py b/etils/enp/numpy_utils.py index 969fe740..ec2f6cab 100644 --- a/etils/enp/numpy_utils.py +++ b/etils/enp/numpy_utils.py @@ -77,31 +77,31 @@ def has_torch(self) -> bool: @property def jax(self): - import jax # pylint: disable=g-import-not-at-top # pytype: disable=import-error + import jax # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return jax @property def jnp(self): - import jax.numpy as jnp # pylint: disable=g-import-not-at-top # pytype: disable=import-error + import jax.numpy as jnp # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return jnp @property def tf(self): - import tensorflow # pylint: disable=g-import-not-at-top # pytype: disable=import-error + import tensorflow # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-source-for-stubs] return tensorflow @property def tnp(self): - import tensorflow.experimental.numpy as tnp # pylint: disable=g-import-not-at-top # pytype: disable=import-error + import tensorflow.experimental.numpy as tnp # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return tnp @property def torch(self): - import torch # pylint: disable=g-import-not-at-top # pytype: disable=import-error + import torch # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] return torch diff --git a/etils/epy/adhoc_utils/curr_args.py b/etils/epy/adhoc_utils/curr_args.py index 57aa6afd..597c1da0 100644 --- a/etils/epy/adhoc_utils/curr_args.py +++ b/etils/epy/adhoc_utils/curr_args.py @@ -70,14 +70,14 @@ def replay_adhoc_ctx(**adhoc_kwargs: Any): match scope: case Scope.COLAB: - from etils import ecolab # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from etils import ecolab # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-module-attribute] return ecolab.adhoc(**adhoc_kwargs) case Scope.BINARY: # Added by LazyModule but not supported by binary_adhoc adhoc_kwargs.pop('collapse_prefix') - from etils.epy.adhoc_utils import binary_import # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from etils.epy.adhoc_utils import binary_import # pylint: disable=g-import-not-at-top return binary_import.binary_adhoc(**adhoc_kwargs) case _: diff --git a/etils/epy/adhoc_utils/module_utils.py b/etils/epy/adhoc_utils/module_utils.py index 04c66139..06fc6ff7 100644 --- a/etils/epy/adhoc_utils/module_utils.py +++ b/etils/epy/adhoc_utils/module_utils.py @@ -114,7 +114,7 @@ def clear_cached_modules( # The typing module has side effect by caching `A[B]` from the old modules # but thankfully they expose the cleanup method. - for cleanup in typing._cleanups: # pytype: disable=module-attr # pylint: disable=protected-access + for cleanup in typing._cleanups: # pylint: disable=protected-access # pyrefly: ignore[missing-attribute] cleanup()