From 1d9b652746a8ed55fbd900edae94bb5bd68dfcdb Mon Sep 17 00:00:00 2001 From: Hana Joo Date: Sun, 2 Aug 2026 06:07:08 -0700 Subject: [PATCH] Add pyrefly suppressions PiperOrigin-RevId: 957930616 --- grain/_src/python/dataset/elastic_iterator_test.py | 8 ++++---- grain/_src/python/dataset/sources/hf_source.py | 2 +- grain/_src/python/dataset/transformations/batch_test.py | 4 ++-- .../python/dataset/transformations/interleave_test.py | 6 +++--- grain/_src/python/dataset/transformations/mix.py | 2 +- grain/_src/python/dataset/transformations/zip.py | 2 +- grain/_src/python/dataset/transformations/zip_test.py | 4 ++-- setup.py | 4 ++-- 8 files changed, 16 insertions(+), 16 deletions(-) diff --git a/grain/_src/python/dataset/elastic_iterator_test.py b/grain/_src/python/dataset/elastic_iterator_test.py index b63e3d9fd..0b7fbd830 100644 --- a/grain/_src/python/dataset/elastic_iterator_test.py +++ b/grain/_src/python/dataset/elastic_iterator_test.py @@ -616,7 +616,7 @@ def _save_elastic_iterators( directory: str, iterators: list[elastic_iterator.ElasticIterator], ): - directory = epath.Path(directory) + directory = epath.Path(directory) # pyrefly: ignore[bad-assignment] checkpoint_handler = handler.CheckpointHandler() for i, iterator in enumerate(iterators): with mock.patch.object( @@ -624,14 +624,14 @@ def _save_elastic_iterators( "get_process_index_and_count", return_value=(i, len(iterators)), ): - checkpoint_handler.save(directory, iterator) + checkpoint_handler.save(directory, iterator) # pyrefly: ignore[bad-argument-type] def _restore_elastic_iterators( self, directory: str, iterators: list[elastic_iterator.ElasticIterator], ): - directory = epath.Path(directory) + directory = epath.Path(directory) # pyrefly: ignore[bad-assignment] checkpoint_handler = handler.CheckpointHandler() for i, iterator in enumerate(iterators): with mock.patch.object( @@ -639,7 +639,7 @@ def _restore_elastic_iterators( "get_process_index_and_count", return_value=(i, len(iterators)), ): - checkpoint_handler.restore(directory, iterator) + checkpoint_handler.restore(directory, iterator) # pyrefly: ignore[bad-argument-type] def test_checkpointing_with_scale_up(self): temp_dir = self.create_tempdir() diff --git a/grain/_src/python/dataset/sources/hf_source.py b/grain/_src/python/dataset/sources/hf_source.py index 9665926c4..baa76374a 100644 --- a/grain/_src/python/dataset/sources/hf_source.py +++ b/grain/_src/python/dataset/sources/hf_source.py @@ -86,7 +86,7 @@ def get_state(self) -> dict[str, Any]: state = {"count_elements_read": self._count_elements_read} _ = self._hf_iter try: - state["hf_state_dict"] = self._hf_ds.state_dict() + state["hf_state_dict"] = self._hf_ds.state_dict() # pyrefly: ignore[bad-assignment] except (AttributeError, NotImplementedError): # Catch when state_dict() is not implemented or not supported. pass diff --git a/grain/_src/python/dataset/transformations/batch_test.py b/grain/_src/python/dataset/transformations/batch_test.py index 3eaeeb8df..6faa323ac 100644 --- a/grain/_src/python/dataset/transformations/batch_test.py +++ b/grain/_src/python/dataset/transformations/batch_test.py @@ -588,7 +588,7 @@ def test_element_spec_custom_batch_fn(self): def test_dynamic_length_and_drop_remainder_when_source_sliced(self): """BatchMapDataset must update length and enforce drop_remainder when its source is sliced.""" - ds = source.SourceMapDataset(list(range(31))) + ds = source.SourceMapDataset(list(range(31))) # pyrefly: ignore[bad-argument-type] batched_ds = batch.BatchMapDataset(ds, batch_size=4, drop_remainder=True) self.assertLen(batched_ds, 7) @@ -604,7 +604,7 @@ def test_dynamic_length_and_drop_remainder_when_source_sliced(self): def test_end_to_end_multithread_prefetch_with_sequential_slice(self): """End-to-end multi-worker pipelines drop remainders cleanly.""" - ds = source.SourceMapDataset(list(range(31))) + ds = source.SourceMapDataset(list(range(31))) # pyrefly: ignore[bad-argument-type] batched_ds = batch.BatchMapDataset(ds, batch_size=4, drop_remainder=True) repeated_ds = batched_ds.repeat(2) iter_ds = repeated_ds.to_iter_dataset() diff --git a/grain/_src/python/dataset/transformations/interleave_test.py b/grain/_src/python/dataset/transformations/interleave_test.py index e35998453..27c6f0868 100644 --- a/grain/_src/python/dataset/transformations/interleave_test.py +++ b/grain/_src/python/dataset/transformations/interleave_test.py @@ -405,7 +405,7 @@ def test_slice_state_management_checkpoints_correctly( ) # Verify it continues from the correct position. - self.assertSequenceEqual(list(it2), expected_remaining) + self.assertSequenceEqual(list(it2), expected_remaining) # pyrefly: ignore[bad-argument-type] @parameterized.named_parameters( dict( @@ -484,7 +484,7 @@ def test_correct_interleave_state_after_setting_shards( ) # Check get_state() internal values. - state = it2.get_state() + state = it2.get_state() # pyrefly: ignore[missing-attribute] self.assertEqual(state["next_index_in_cycle"], 0) self.assertEqual( state["next_index_in_datasets"], expected_next_index_in_datasets @@ -520,7 +520,7 @@ def test_setting_shard_state_with_exhausted_states(self): ) # Check get_state() internal values. - state = it.get_state() + state = it.get_state() # pyrefly: ignore[missing-attribute] self.assertEqual(state["next_index_in_cycle"], 0) self.assertEqual(state["next_index_in_datasets"], 3) self.assertEqual(state["iterators_in_use_indices"], [2, 0]) diff --git a/grain/_src/python/dataset/transformations/mix.py b/grain/_src/python/dataset/transformations/mix.py index 192648980..7c4d9a856 100644 --- a/grain/_src/python/dataset/transformations/mix.py +++ b/grain/_src/python/dataset/transformations/mix.py @@ -356,7 +356,7 @@ def __iter__(self) -> dataset.DatasetIterator[T]: def set_slice(self, sl: slice, sequential_slice: bool = False) -> None: for parent in self._parents: - dataset.set_slice(parent, sl, sequential_slice) + dataset.set_slice(parent, sl, sequential_slice) # pyrefly: ignore[bad-argument-type] def __str__(self) -> str: return ( diff --git a/grain/_src/python/dataset/transformations/zip.py b/grain/_src/python/dataset/transformations/zip.py index 12691b546..511a44c17 100644 --- a/grain/_src/python/dataset/transformations/zip.py +++ b/grain/_src/python/dataset/transformations/zip.py @@ -77,7 +77,7 @@ def __iter__(self) -> dataset.DatasetIterator[T]: def set_slice(self, sl: slice, sequential_slice: bool = False) -> None: del sequential_slice for parent in self._parents: - dataset.set_slice(parent, sl) + dataset.set_slice(parent, sl) # pyrefly: ignore[bad-argument-type] def __str__(self) -> str: return f"ZipIterDataset(parents={self._parents}, strict={self._strict})" diff --git a/grain/_src/python/dataset/transformations/zip_test.py b/grain/_src/python/dataset/transformations/zip_test.py index 717cd35fd..2fd5578b9 100644 --- a/grain/_src/python/dataset/transformations/zip_test.py +++ b/grain/_src/python/dataset/transformations/zip_test.py @@ -295,12 +295,12 @@ def __init__(self, parent): self._parents = [parent] def __iter__(self): - return WrapperDatasetIterator(self._parents[0].__iter__()) + return WrapperDatasetIterator(self._parents[0].__iter__()) # pyrefly: ignore[missing-attribute] class WrapperDatasetIterator(dataset.DatasetIterator): def __init__(self, parent_iter): - self._parents = [parent_iter] + self._parents = [parent_iter] # pyrefly: ignore[bad-assignment] self._ctx = parent_iter._ctx def __next__(self): diff --git a/setup.py b/setup.py index 9ae3781fe..ca6146d75 100644 --- a/setup.py +++ b/setup.py @@ -47,7 +47,7 @@ def finalize_options(self): pass def run(self): - from grpc_tools import protoc # pylint: disable=g-import-not-at-top + from grpc_tools import protoc # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] root_dir = os.path.dirname(os.path.abspath(__file__)) @@ -77,7 +77,7 @@ def run(self): self.run_command("generate_protos") super().run() - from pybind11.setup_helpers import Pybind11Extension # pylint: disable=g-import-not-at-top + from pybind11.setup_helpers import Pybind11Extension # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] ext_modules = [ Pybind11Extension(