Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 7 additions & 7 deletions etils/enp/array_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,34 +165,34 @@ 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


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,
'SharedMemoryArrayMetadata',
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


Expand All @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion etils/enp/checking.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
16 changes: 8 additions & 8 deletions etils/enp/checking_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
2 changes: 1 addition & 1 deletion etils/enp/compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion etils/enp/interp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
10 changes: 5 additions & 5 deletions etils/enp/numpy_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
4 changes: 2 additions & 2 deletions etils/epy/adhoc_utils/curr_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 _:
Expand Down
2 changes: 1 addition & 1 deletion etils/epy/adhoc_utils/module_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()


Expand Down
Loading