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
2 changes: 1 addition & 1 deletion etils/eapp/logging_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ class TqdmStream:
"""

def write(self, x: str) -> None:
import tqdm # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import tqdm # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-source-for-stubs]

tqdm.tqdm.write(x, end='')

Expand Down
8 changes: 4 additions & 4 deletions etils/ecolab/inspects/nodes.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ def from_obj(cls, obj: object, *, name: str = '') -> ObjectNode:
else:
raise TypeError(f'Unexpected object {obj!r}.')

return sub_cls(obj=obj, name=name) # pytype: disable=wrong-arg-types
return sub_cls(obj=obj, name=name)

@property
def header_html(self) -> str:
Expand Down Expand Up @@ -252,7 +252,7 @@ def is_leaf(self) -> bool:


@dataclasses.dataclass
class BuiltinNode(ObjectNode[Union[int, float, bool, str, bytes, None]]): # pytype: disable=bad-concrete-type
class BuiltinNode(ObjectNode[Union[int, float, bool, str, bytes, None]]):
"""`int`, `float`, `bytes`, `str`,..."""

MATCH_TYPES = (type(None), int, float, bool, str, bytes, type(...))
Expand All @@ -267,7 +267,7 @@ def is_leaf(self) -> bool:


@dataclasses.dataclass
class MappingNode(ObjectNode[collections.abc.Mapping]): # pytype: disable=bad-concrete-type
class MappingNode(ObjectNode[collections.abc.Mapping]):
"""`dict` like."""

MATCH_TYPES = (
Expand All @@ -286,7 +286,7 @@ def children(self) -> list[Node]:


@dataclasses.dataclass
class SetNode(ObjectNode[collections.abc.Set]): # pytype: disable=bad-concrete-type
class SetNode(ObjectNode[collections.abc.Set]):
"""`set` like."""

MATCH_TYPES = (set, frozenset, collections.abc.Set)
Expand Down
2 changes: 1 addition & 1 deletion etils/ecolab/pyjs_com/py_js_com.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ def wrap_output(self, out):
return IPython.display.JSON(out)

def register_fn(self, fn: _Fn) -> None:
from google.colab import output # pylint: disable=g-import-not-at-top # pytype: disable=import-error
from google.colab import output # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

# TODO(epot): Fragile if multiple functions have the same name. How to
# specify namespace ?
Expand Down
58 changes: 29 additions & 29 deletions etils/enp/array_types/dtypes_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,32 +64,32 @@ class _DTypeTestItem:
name='AnyDType',
cls_name='Array',
array_items=[
_ArrayItem(True, np.bool_), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([True, False], np.bool_), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(True, np.bool_), # pyrefly: ignore[bad-argument-type]
_ArrayItem([True, False], np.bool_), # pyrefly: ignore[bad-argument-type]
# Numpy default are plaform dependent (int64 except on windows)
# _ArrayItem(1, np.int32),
# _ArrayItem(1., np.float32),
# _ArrayItem([1], np.int32),
# _ArrayItem([1.], np.float32),
_ArrayItem(np.array(1, dtype=np.uint8), np.uint8), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array(1, dtype=np.uint8), np.uint8), # pyrefly: ignore[bad-argument-type]
],
),
_DTypeTestItem(
name='AnyFloat',
cls_name='FloatArray',
array_items=[
# Bool, int,... casted to float
_ArrayItem(True, np.float32, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1, np.float32, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1.0, np.float32), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1], np.float32, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1.0], np.float32), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem( # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(True, np.float32, iscast=True), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1, np.float32, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1.0, np.float32), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1], np.float32, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1.0], np.float32), # pyrefly: ignore[bad-argument-type]
_ArrayItem(
np.array([True], dtype=np.bool_), np.float32, iscast=True # pyrefly: ignore[bad-argument-type]
),
_ArrayItem(np.array(1, dtype=np.uint8), np.float32, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array(1, dtype=np.uint8), np.float32, iscast=True), # pyrefly: ignore[bad-argument-type]
# Float values not casted
_ArrayItem(np.array(1, dtype=np.float16), np.float16), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array(1, dtype=np.float16), np.float16), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array(1, dtype=jnp.bfloat16), jnp.bfloat16), # pyrefly: ignore[bad-argument-type]
],
),
Expand All @@ -98,16 +98,16 @@ class _DTypeTestItem:
cls_name='IntArray',
array_items=[
# Casted to int
_ArrayItem(True, np.int32, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1, np.int32), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1.0, np.int32, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1], np.int32), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1.0], np.int32, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(True, np.int32, iscast=True), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1, np.int32), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1.0, np.int32, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1], np.int32), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1.0], np.int32, iscast=False), # pyrefly: ignore[bad-argument-type]
# Int values not casted
_ArrayItem(np.array([True], dtype=np.bool_), np.int32, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array(1, dtype=np.uint8), np.uint8), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array([True], dtype=np.bool_), np.int32, iscast=True), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array(1, dtype=np.uint8), np.uint8), # pyrefly: ignore[bad-argument-type]
# Float values casted
_ArrayItem(np.array(1, dtype=np.float16), np.int32, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array(1, dtype=np.float16), np.int32, iscast=True), # pyrefly: ignore[bad-argument-type]
],
),
# `BoolArray` -> `bool_`
Expand All @@ -125,22 +125,22 @@ class _DTypeTestItem:
name='bool',
cls_name='bool_',
array_items=[
_ArrayItem(True, np.bool_), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1, np.bool_, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1.0], np.bool_, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array([True], dtype=np.bool_), np.bool_), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array([1], dtype=np.uint8), np.bool_, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(True, np.bool_), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1, np.bool_, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1.0], np.bool_, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array([True], dtype=np.bool_), np.bool_), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array([1], dtype=np.uint8), np.bool_, iscast=True), # pyrefly: ignore[bad-argument-type]
],
),
_DTypeTestItem(
name='uint8',
cls_name='ui8',
array_items=[
_ArrayItem(True, np.uint8, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(1, np.uint8), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem([1.0], np.uint8, iscast=False), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array([True], dtype=np.bool_), np.uint8, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(np.array([1], dtype=np.int32), np.uint8, iscast=True), # pytype: disable=wrong-arg-types # numpy-scalars
_ArrayItem(True, np.uint8, iscast=True), # pyrefly: ignore[bad-argument-type]
_ArrayItem(1, np.uint8), # pyrefly: ignore[bad-argument-type]
_ArrayItem([1.0], np.uint8, iscast=False), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array([True], dtype=np.bool_), np.uint8, iscast=True), # pyrefly: ignore[bad-argument-type]
_ArrayItem(np.array([1], dtype=np.int32), np.uint8, iscast=True), # pyrefly: ignore[bad-argument-type]
],
),
_DTypeTestItem(
Expand Down
2 changes: 1 addition & 1 deletion etils/enp/array_types/typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,7 @@ def __new__( # pylint: disable=bad-mcs-classmethod-argument

def __init__(cls, shape: Optional[ShapeSpec], dtype: Optional[_DType]):
del shape, dtype
super().__init__(cls, cls.__name__, (cls,), {}) # pytype: disable=wrong-arg-count
super().__init__(cls, cls.__name__, (cls,), {}) # pyrefly: ignore[no-matching-overload]

def __getitem__(cls, shape: _ShapeSpecInput) -> 'ArrayAliasMeta':
if shape is None: # Normalize 'Array[None]'
Expand Down
10 changes: 5 additions & 5 deletions etils/etree/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@ class Jax(Backend):
"""`jax.tree_util` backend."""

def import_module(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.tree_util

Expand Down Expand Up @@ -134,7 +134,7 @@ class Optree(Jax):
"""`optree` backend."""

def import_module(self):
import optree # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import optree # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

return optree

Expand Down Expand Up @@ -172,7 +172,7 @@ class Nest(Backend):
"""`tf.nest` backend."""

def import_module(self):
import tensorflow as tf # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import tensorflow as tf # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-source-for-stubs]

return tf.nest

Expand Down Expand Up @@ -229,7 +229,7 @@ def map(self, map_fn, *trees, is_leaf=None):
new_tree.update(new_items)
return new_tree
else:
return type(tree0)(new_items) # pyrefly: ignore[bad-argument-count, bad-instantiation]
return type(tree0)(new_items) # pyrefly: ignore[bad-argument-count]
else: # leaf
return map_fn(*trees)

Expand Down Expand Up @@ -275,7 +275,7 @@ def _unflatten(self, structure, flat_iter):
new_tree.update(new_items)
return new_tree
else:
return type(structure)(new_items) # pyrefly: ignore[bad-argument-count, bad-instantiation]
return type(structure)(new_items) # pyrefly: ignore[bad-argument-count]
else: # leaf
return next(flat_iter)

Expand Down
2 changes: 1 addition & 1 deletion etils/etree/typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,6 @@

# TODO(b/202712189): Add `Any` to the `Union`
# Tree can be arbitrary `chex.dataclass`, jax structure,... so has to be Any
Tree = Union[_T, list['Tree'], tuple['Tree', ...], dict[str, 'Tree']] # pytype: disable=not-supported-yet
Tree = Union[_T, list['Tree'], tuple['Tree', ...], dict[str, 'Tree']]

LeafFn = Callable[[Any], bool]
Loading