From c0a9f039d8b8d128f69b626d391e0d5e06aa46f4 Mon Sep 17 00:00:00 2001 From: Oleh Prypin Date: Tue, 6 Oct 2026 07:59:01 -0700 Subject: [PATCH] Replace `# pytype: disable` suppressions with `# pyrefly: ignore` PiperOrigin-RevId: 994394434 --- etils/eapp/logging_utils.py | 2 +- etils/ecolab/inspects/nodes.py | 8 ++-- etils/ecolab/pyjs_com/py_js_com.py | 2 +- etils/enp/array_types/dtypes_test.py | 58 ++++++++++++++-------------- etils/enp/array_types/typing.py | 2 +- etils/etree/backend.py | 10 ++--- etils/etree/typing.py | 2 +- 7 files changed, 42 insertions(+), 42 deletions(-) diff --git a/etils/eapp/logging_utils.py b/etils/eapp/logging_utils.py index ee29f6cb..40466ba6 100644 --- a/etils/eapp/logging_utils.py +++ b/etils/eapp/logging_utils.py @@ -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='') diff --git a/etils/ecolab/inspects/nodes.py b/etils/ecolab/inspects/nodes.py index 689a439d..d6ad226c 100644 --- a/etils/ecolab/inspects/nodes.py +++ b/etils/ecolab/inspects/nodes.py @@ -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: @@ -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(...)) @@ -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 = ( @@ -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) diff --git a/etils/ecolab/pyjs_com/py_js_com.py b/etils/ecolab/pyjs_com/py_js_com.py index 74f401f4..4c3db9f3 100644 --- a/etils/ecolab/pyjs_com/py_js_com.py +++ b/etils/ecolab/pyjs_com/py_js_com.py @@ -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 ? diff --git a/etils/enp/array_types/dtypes_test.py b/etils/enp/array_types/dtypes_test.py index e939311f..fec4a928 100644 --- a/etils/enp/array_types/dtypes_test.py +++ b/etils/enp/array_types/dtypes_test.py @@ -64,14 +64,14 @@ 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( @@ -79,17 +79,17 @@ class _DTypeTestItem: 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] ], ), @@ -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_` @@ -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( diff --git a/etils/enp/array_types/typing.py b/etils/enp/array_types/typing.py index 90e87cd3..ecf6c108 100644 --- a/etils/enp/array_types/typing.py +++ b/etils/enp/array_types/typing.py @@ -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]' diff --git a/etils/etree/backend.py b/etils/etree/backend.py index 16a1fd5e..b3cf45f8 100644 --- a/etils/etree/backend.py +++ b/etils/etree/backend.py @@ -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 @@ -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 @@ -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 @@ -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) @@ -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) diff --git a/etils/etree/typing.py b/etils/etree/typing.py index 1b3dee79..245b7657 100644 --- a/etils/etree/typing.py +++ b/etils/etree/typing.py @@ -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]