diff --git a/airio/_src/core/data_sources.py b/airio/_src/core/data_sources.py index 288563f..6edba75 100644 --- a/airio/_src/core/data_sources.py +++ b/airio/_src/core/data_sources.py @@ -22,7 +22,7 @@ class DataSource(Protocol): """Interface for data sources wrappers with multiple splits support.""" - splits: Iterable[str] = None + splits: Iterable[str] = None # pyrefly: ignore[bad-assignment] def get_data_source(self, split: str): ... diff --git a/airio/_src/core/dataset_iterators.py b/airio/_src/core/dataset_iterators.py index f4deb22..27d9243 100644 --- a/airio/_src/core/dataset_iterators.py +++ b/airio/_src/core/dataset_iterators.py @@ -24,7 +24,7 @@ class AirIODatasetIterator(clu_dataset_iterator.DatasetIterator): """Wrapper iterator for AirIO.""" - _iterator: collections.abc.Iterator[Any] = None + _iterator: collections.abc.Iterator[Any] = None # pyrefly: ignore[bad-assignment] def __next__(self) -> clu_dataset_iterator.Element: raise NotImplementedError() diff --git a/airio/_src/core/dataset_providers.py b/airio/_src/core/dataset_providers.py index aa12223..2cb890e 100644 --- a/airio/_src/core/dataset_providers.py +++ b/airio/_src/core/dataset_providers.py @@ -43,12 +43,12 @@ class ShardInfo: class DatasetProviderBase(Protocol): """Abstract base for classes that provide a dataset.""" - splits: Iterable[str] = None + splits: Iterable[str] = None # pyrefly: ignore[bad-assignment] def get_dataset( self, sequence_lengths: Mapping[str, int] | None = None, - split: str = tfds.Split.TRAIN, + split: str = tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute] runtime_preprocessors: Sequence[grain.Transformation] | None = None, batch_size: int | None = None, shuffle: bool = True, @@ -114,7 +114,7 @@ def __init__( all_tasks = [t for t in tasks if isinstance(t, Task)] all_mixtures = [m for m in tasks if isinstance(m, Mixture)] sub_tasks = [mix.leaf_tasks for mix in all_mixtures] - leaf_tasks = sum(sub_tasks, all_tasks) + leaf_tasks = sum(sub_tasks, all_tasks) # pyrefly: ignore[no-matching-overload] duplicate_tasks = [ t for t, c in collections.Counter(leaf_tasks).items() if c > 1 ] @@ -128,7 +128,7 @@ def __init__( self._proportions = dict(zip(tasks, proportions)) def num_input_examples(self, split: str) -> int | None: - return sum( + return sum( # pyrefly: ignore[no-matching-overload] t.num_input_examples(split) for t in self.tasks_or_mixtures if split in t.splits @@ -164,14 +164,14 @@ def leaf_tasks(self) -> Sequence[Task]: tasks = [t for t in all_ if isinstance(t, Task)] mixtures = [m for m in all_ if isinstance(m, Mixture)] sub_tasks = [mix.leaf_tasks for mix in mixtures] - return sum(sub_tasks, tasks) + return sum(sub_tasks, tasks) # pyrefly: ignore[no-matching-overload] @property def total_proportion(self) -> float: return sum(self._proportions.values()) @property - def splits(self) -> Sequence[str]: + def splits(self) -> Sequence[str]: # pyrefly: ignore[bad-override] splits = set() for task in self.tasks_or_mixtures: splits.update(task.splits) @@ -218,7 +218,7 @@ def build(self) -> Task: if self._preprocessors is None: raise ValueError("Preprocessors have not been set on this task builder.") - return Task( + return Task( # pyrefly: ignore[bad-instantiation] name=self._task_name, source=self._source, preprocessors=self._preprocessors, @@ -248,7 +248,7 @@ def from_task(cls, task: Task) -> "TaskBuilder": Args: task: Existing task object. """ - return TaskBuilder( + return TaskBuilder( # pyrefly: ignore[bad-instantiation] task_name=task.name, source=task.source, preprocessors=task.get_preprocessors(), diff --git a/airio/_src/core/preprocessors.py b/airio/_src/core/preprocessors.py index 17b287e..820dcb7 100644 --- a/airio/_src/core/preprocessors.py +++ b/airio/_src/core/preprocessors.py @@ -60,7 +60,7 @@ def replace(self, **kwargs): @dataclasses.dataclass @typing.runtime_checkable -class MapFnTransform(Protocol): +class MapFnTransform(Protocol): # pyrefly: ignore[bad-class-definition] """Transform to represent AirIO map preprocessors. Attrs: @@ -84,7 +84,7 @@ def map(self, element): @dataclasses.dataclass @typing.runtime_checkable -class RandomMapFnTransform(Protocol): +class RandomMapFnTransform(Protocol): # pyrefly: ignore[bad-class-definition] """Transform to represent AirIO random map preprocessors. Attrs: @@ -108,7 +108,7 @@ def random_map(self, element, rng: np.random.Generator): @dataclasses.dataclass @typing.runtime_checkable -class FilterFnTransform(Protocol): +class FilterFnTransform(Protocol): # pyrefly: ignore[bad-class-definition] """Transform to represent AirIO filter preprocessors. Attrs: diff --git a/airio/_src/core/test_utils.py b/airio/_src/core/test_utils.py index ec441a6..11ee428 100644 --- a/airio/_src/core/test_utils.py +++ b/airio/_src/core/test_utils.py @@ -35,7 +35,7 @@ def assert_datasets_equal( """ if not isinstance(expected, list): - expected = [expected] + expected = [expected] # pyrefly: ignore[bad-assignment] actual = list(dataset) absltest.TestCase().assertEqual(len(actual), len(expected)) @@ -79,4 +79,4 @@ def create_airio_injected_runtime_args( "batch_size": batch_size, } args = {k: provided[k] if provided[k] else defaults[k] for k in defaults} - return preprocessors.AirIOInjectedRuntimeArgs(**args) + return preprocessors.AirIOInjectedRuntimeArgs(**args) # pyrefly: ignore[bad-argument-type] diff --git a/airio/_src/core/tokenizer.py b/airio/_src/core/tokenizer.py index baf917c..42b8780 100644 --- a/airio/_src/core/tokenizer.py +++ b/airio/_src/core/tokenizer.py @@ -38,7 +38,7 @@ def vocabulary(self) -> vocabularies.Vocabulary: @typing.runtime_checkable @dataclasses.dataclass(frozen=True) -class Tokenizer(Generic[Inp, Out], Protocol): +class Tokenizer(Generic[Inp, Out], Protocol): # pyrefly: ignore[bad-class-definition] """Tokenizer class for AirIO tasks/mixtures.""" tokenizer_configs: Mapping[str, TokenizerConfig] diff --git a/airio/_src/core/vocabularies.py b/airio/_src/core/vocabularies.py index 01abeee..00b6110 100644 --- a/airio/_src/core/vocabularies.py +++ b/airio/_src/core/vocabularies.py @@ -237,7 +237,7 @@ def _load_model( return cls._ModelContext(tokenizer=tokenizer, sp_model=sp_model) @property - def pad_id(self) -> int | None: + def pad_id(self) -> int | None: # pyrefly: ignore[bad-override] return PAD_ID @property @@ -287,7 +287,7 @@ def __eq__(self, other): if not isinstance(other, SentencePieceVocabulary): return False try: - their_md5 = hashlib.md5(other.sp_model).hexdigest() + their_md5 = hashlib.md5(other.sp_model).hexdigest() # pyrefly: ignore[bad-argument-type] # If other has no sp_model attribute, we can't test for equality except AttributeError: return False @@ -298,6 +298,7 @@ def __eq__(self, other): def __str__(self) -> str: return ( + # pyrefly: ignore[bad-argument-type] f"SentencePieceVocabulary(file={self.sentencepiece_model_file}, " f"extra_ids={self._extra_ids}, " f"spm_md5={hashlib.md5(self.sp_model).hexdigest()})" @@ -344,7 +345,7 @@ def eos_id(self) -> int | None: return None @property - def pad_id(self) -> int | None: + def pad_id(self) -> int | None: # pyrefly: ignore[bad-override] return PAD_ID @property diff --git a/airio/_src/pygrain/common/feature_converters.py b/airio/_src/pygrain/common/feature_converters.py index c43d27c..a3be8bb 100644 --- a/airio/_src/pygrain/common/feature_converters.py +++ b/airio/_src/pygrain/common/feature_converters.py @@ -101,22 +101,22 @@ def get_t5x_enc_dec_feature_converter_preprocessors( ) pack_prep.append( preprocessors.LazyIterTransform( - packer, update_runtime_args=packer.update_runtime_args + packer, update_runtime_args=packer.update_runtime_args # pyrefly: ignore[bad-argument-type] ) ) return ( [ preprocessors.MapFnTransform( - common_preprocessors.remove_features_not_in_sequence_lengths + common_preprocessors.remove_features_not_in_sequence_lengths # pyrefly: ignore[bad-argument-count] ), - preprocessors.MapFnTransform(common_preprocessors.trim), + preprocessors.MapFnTransform(common_preprocessors.trim), # pyrefly: ignore[bad-argument-count] ] + pack_prep + [ - preprocessors.MapFnTransform(pad), + preprocessors.MapFnTransform(pad), # pyrefly: ignore[bad-argument-count] preprocessors.MapFnTransform( - convert_features, - update_runtime_args=update_runtime_args, + convert_features, # pyrefly: ignore[bad-argument-count] + update_runtime_args=update_runtime_args, # pyrefly: ignore[unexpected-keyword] ), ] ) @@ -167,22 +167,22 @@ def get_t5x_lm_feature_converter_preprocessors( ) packer_prep.append( preprocessors.LazyIterTransform( - packer, update_runtime_args=packer.update_runtime_args + packer, update_runtime_args=packer.update_runtime_args # pyrefly: ignore[bad-argument-type] ) ) return ( [ preprocessors.MapFnTransform( - common_preprocessors.remove_features_not_in_sequence_lengths + common_preprocessors.remove_features_not_in_sequence_lengths # pyrefly: ignore[bad-argument-count] ), - preprocessors.MapFnTransform(common_preprocessors.trim), + preprocessors.MapFnTransform(common_preprocessors.trim), # pyrefly: ignore[bad-argument-count] ] + packer_prep + [ - preprocessors.MapFnTransform(pad), + preprocessors.MapFnTransform(pad), # pyrefly: ignore[bad-argument-count] preprocessors.MapFnTransform( - convert_features, - update_runtime_args=update_runtime_args, + convert_features, # pyrefly: ignore[bad-argument-count] + update_runtime_args=update_runtime_args, # pyrefly: ignore[unexpected-keyword] ), ] ) @@ -258,9 +258,9 @@ def swap_inputs_width(ex: dict[str, np.ndarray], old_val: int, new_val: int): preps = [ preprocessors.MapFnTransform( - concat_and_add_masks, update_runtime_args=concat_task_feature_lengths + concat_and_add_masks, update_runtime_args=concat_task_feature_lengths # pyrefly: ignore[bad-argument-count, unexpected-keyword] ), - preprocessors.MapFnTransform(replace_0s), + preprocessors.MapFnTransform(replace_0s), # pyrefly: ignore[bad-argument-count] ] if pack: packer = ( @@ -269,15 +269,15 @@ def swap_inputs_width(ex: dict[str, np.ndarray], old_val: int, new_val: int): else packing.SingleBinTruePackIterPreprocessor ) packer_prep = preprocessors.LazyIterTransform( - packer, update_runtime_args=packer.update_runtime_args + packer, update_runtime_args=packer.update_runtime_args # pyrefly: ignore[bad-argument-type] ) - preps.append(packer_prep) + preps.append(packer_prep) # pyrefly: ignore[bad-argument-type] preps.extend([ - preprocessors.MapFnTransform(common_preprocessors.trim), - preprocessors.MapFnTransform(pad), - preprocessors.MapFnTransform(restore_0s), + preprocessors.MapFnTransform(common_preprocessors.trim), # pyrefly: ignore[bad-argument-count] + preprocessors.MapFnTransform(pad), # pyrefly: ignore[bad-argument-count] + preprocessors.MapFnTransform(restore_0s), # pyrefly: ignore[bad-argument-count] preprocessors.MapFnTransform( - convert_features, update_runtime_args=update_runtime_args + convert_features, update_runtime_args=update_runtime_args # pyrefly: ignore[bad-argument-count, unexpected-keyword] ), ]) return preps diff --git a/airio/_src/pygrain/common/packing.py b/airio/_src/pygrain/common/packing.py index da8ce84..3aa931a 100644 --- a/airio/_src/pygrain/common/packing.py +++ b/airio/_src/pygrain/common/packing.py @@ -221,7 +221,7 @@ def parent(self): def __len__(self): return len(self.parent) - def __getitem__(self, index: slice): + def __getitem__(self, index: slice): # pyrefly: ignore[bad-override] if isinstance(index, slice): return self.slice(index) return self._packed_ds[index] @@ -350,7 +350,7 @@ def get_state(self): def set_state(self, state): self._parent_iter.set_state(state["parent"]) - self._packer = self._packer_type.from_dict(state["packer"]) + self._packer = self._packer_type.from_dict(state["packer"]) # pyrefly: ignore[bad-argument-type] self._packed_examples = collections.deque[PyTree[np.ndarray]]( [load_np_tree(t) for t in state["packed_examples"]] ) @@ -514,7 +514,7 @@ def fit_example(self, ex: PyTree[np.ndarray]) -> Sequence[PyTree[np.ndarray]]: # Add if example fits an existing partially packed example; check if # resulting partially packed example becomes fully packed fits = False - fully_packed: PartiallyPackedExample = None + fully_packed: PartiallyPackedExample = None # pyrefly: ignore[bad-assignment] fully_packed_idx = None for idx, partially_packed in enumerate(self._partially_packed_examples): if partially_packed.example_fits(flat_ex): @@ -532,7 +532,7 @@ def fit_example(self, ex: PyTree[np.ndarray]) -> Sequence[PyTree[np.ndarray]]: fully_packed.pack(), length_struct=self.feature_lengths, ) - del self._partially_packed_examples[fully_packed_idx] + del self._partially_packed_examples[fully_packed_idx] # pyrefly: ignore[unsupported-operation] # self._partially_packed_examples.remove(fully_packed) packed_examples.append(packed) @@ -615,7 +615,7 @@ def __init__( ): self._feature_lengths = feature_lengths self._flat_feature_lengths = flatten(feature_lengths) - self._partially_packed_example: PartiallyPackedExample = None + self._partially_packed_example: PartiallyPackedExample = None # pyrefly: ignore[bad-assignment] if feature_lengths: self._partially_packed_example = PartiallyPackedExample( copy.copy(self._flat_feature_lengths) diff --git a/airio/_src/pygrain/data_sources.py b/airio/_src/pygrain/data_sources.py index cabd1f4..b81658c 100644 --- a/airio/_src/pygrain/data_sources.py +++ b/airio/_src/pygrain/data_sources.py @@ -39,7 +39,7 @@ def __init__( self.splits = frozenset(self._split_to_filepattern.keys()) self._sources = { - split: grain.ArrayRecordDataSource(self._split_to_filepattern[split]) + split: grain.ArrayRecordDataSource(self._split_to_filepattern[split]) # pyrefly: ignore[bad-argument-type] for split in self.splits } @@ -119,7 +119,7 @@ def __init__( self.splits = frozenset(self._split_to_filepattern.keys()) self._sources = {} for split in self.splits: - json_data = json.load(Open(self._split_to_filepattern[split])) + json_data = json.load(Open(self._split_to_filepattern[split])) # pyrefly: ignore[bad-argument-type] json_data = [json.dumps(d) for d in json_data] self._sources[split] = grain.InMemoryDataSource(elements=json_data) diff --git a/airio/_src/pygrain/dataset_iterators.py b/airio/_src/pygrain/dataset_iterators.py index 2e70394..8db3128 100644 --- a/airio/_src/pygrain/dataset_iterators.py +++ b/airio/_src/pygrain/dataset_iterators.py @@ -98,13 +98,13 @@ def peek_async( def get_state(self) -> Mapping[str, Any]: if self._state_as_dict: - return self._iterator.get_state() - return json.loads(self._iterator.get_state().decode()) + return self._iterator.get_state() # pyrefly: ignore[missing-attribute] + return json.loads(self._iterator.get_state().decode()) # pyrefly: ignore[missing-attribute] def set_state(self, state: Mapping[str, Any]) -> None: if not self._state_as_dict: - state = json.dumps(state, indent=4).encode() - self._iterator.set_state(state) + state = json.dumps(state, indent=4).encode() # pyrefly: ignore[bad-assignment] + self._iterator.set_state(state) # pyrefly: ignore[missing-attribute] def save(self, filename: epath.PathLike): filename = epath.Path(filename) diff --git a/airio/_src/pygrain/dataset_providers.py b/airio/_src/pygrain/dataset_providers.py index 188c55d..ed664d4 100644 --- a/airio/_src/pygrain/dataset_providers.py +++ b/airio/_src/pygrain/dataset_providers.py @@ -63,7 +63,7 @@ def _switch_to_lazy_dataset( # lazy_dataset. preps = self.get_preprocessors() if runtime_preprocessors: - preps.extend(runtime_preprocessors) + preps.extend(runtime_preprocessors) # pyrefly: ignore[bad-argument-type] for preprocessor in preps: if not isinstance(preprocessor, grain.Transformation): return True @@ -109,16 +109,16 @@ def get_lazy_dataset( # Step 3: Run preprocessors and shuffle each epoch (if needed) preps = self._preprocessors updated_runtime_args = core_preprocessors_lib.AirIOInjectedRuntimeArgs( - sequence_lengths=sequence_lengths, + sequence_lengths=sequence_lengths, # pyrefly: ignore[bad-argument-type] split=split, batch_size=batch_size, ) preprocessed_dss = [] has_none_elems = False - next_epoch_rng = jax.random.key(seed) + next_epoch_rng = jax.random.key(seed) # pyrefly: ignore[bad-argument-type] for ds in dss: ds_runtime_args = core_preprocessors_lib.AirIOInjectedRuntimeArgs( - sequence_lengths=sequence_lengths, + sequence_lengths=sequence_lengths, # pyrefly: ignore[bad-argument-type] split=split, batch_size=batch_size, ) @@ -169,10 +169,10 @@ def get_lazy_dataset( return ds # TODO(sahildua): Add logging. - def get_dataset( + def get_dataset( # pyrefly: ignore[bad-override] self, sequence_lengths: Mapping[str, int] | None = None, - split: str = tfds.Split.TRAIN, + split: str = tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute] runtime_preprocessors: ( Sequence[preprocessors_lib.PyGrainAirIOPreprocessor] | None ) = None, @@ -215,7 +215,7 @@ def get_dataset( ) sampler = grain.IndexSampler( - num_records=self.num_input_examples(split=split), + num_records=self.num_input_examples(split=split), # pyrefly: ignore[bad-argument-type] shard_options=shard_options, shuffle=shuffle, num_epochs=num_epochs, @@ -226,13 +226,13 @@ def get_dataset( ops = self.get_preprocessors() if runtime_preprocessors: - ops.extend(runtime_preprocessors) + ops.extend(runtime_preprocessors) # pyrefly: ignore[bad-argument-type] if batch_size: ops.append(grain.Batch(batch_size=batch_size, drop_remainder=False)) # Add runtime args runtime_args = core_preprocessors_lib.AirIOInjectedRuntimeArgs( - sequence_lengths=sequence_lengths, + sequence_lengths=sequence_lengths, # pyrefly: ignore[bad-argument-type] split=split, batch_size=batch_size, ) @@ -291,7 +291,7 @@ def get_dataset_by_step( self, num_records: int = DEFAULT_NUM_RECORDS_TO_INSPECT, sequence_lengths: Mapping[str, int] | None = None, - split: str = tfds.Split.TRAIN, + split: str = tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute] runtime_preprocessors: ( Sequence[preprocessors_lib.PyGrainAirIOPreprocessor] | None ) = None, @@ -341,7 +341,7 @@ def get_dataset_by_step( all_ops = self.get_preprocessors() if runtime_preprocessors: - all_ops.extend(runtime_preprocessors) + all_ops.extend(runtime_preprocessors) # pyrefly: ignore[bad-argument-type] if batch_size: all_ops.append(grain.Batch(batch_size=batch_size, drop_remainder=False)) @@ -354,7 +354,7 @@ def get_dataset_by_step( # Apply all transformations, one by one. runtime_args = core_preprocessors_lib.AirIOInjectedRuntimeArgs( - sequence_lengths=sequence_lengths, + sequence_lengths=sequence_lengths, # pyrefly: ignore[bad-argument-type] split=split, batch_size=batch_size, ) @@ -380,7 +380,7 @@ def get_updated_runtime_args( """Returns updated runtime args based on preprocessors and feature converter.""" preps = self._preprocessors if runtime_preprocessors: - preps.extend(runtime_preprocessors) + preps.extend(runtime_preprocessors) # pyrefly: ignore[bad-argument-type] for prep in preps: transform = preprocessors_lib.LazyDatasetTransform(prep) runtime_args = transform.get_updated_runtime_args(runtime_args) @@ -420,7 +420,7 @@ def __init__( def get_lazy_dataset( self, sequence_lengths: Mapping[str, int] | None = None, - split: str = tfds.Split.TRAIN, + split: str = tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute] runtime_preprocessors: ( Sequence[preprocessors_lib.PyGrainAirIOPreprocessor] | None ) = None, @@ -495,7 +495,7 @@ def get_lazy_dataset( # args must match, or mixing won't work (compute all updated runtime args # and add a check here in the future if helpful). runtime_args = core_preprocessors_lib.AirIOInjectedRuntimeArgs( - sequence_lengths=sequence_lengths, + sequence_lengths=sequence_lengths, # pyrefly: ignore[bad-argument-type] split=split, batch_size=batch_size, ) @@ -503,7 +503,7 @@ def get_lazy_dataset( runtime_args = self.leaf_tasks[0].get_updated_runtime_args( runtime_args, runtime_preprocessors=None ) - base_rng = jax.random.key(seed) + base_rng = jax.random.key(seed) # pyrefly: ignore[bad-argument-type] post_mix_rng, _ = jax.random.split(base_rng) ds, _, _ = _apply_preprocessors_to_lazy_dataset( ds, @@ -519,10 +519,10 @@ def get_lazy_dataset( ds = ds.repeat(num_epochs=None) # pytype: disable=attribute-error return ds - def get_dataset( + def get_dataset( # pyrefly: ignore[bad-override] self, sequence_lengths: Mapping[str, int] | None = None, - split: str = tfds.Split.TRAIN, + split: str = tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute] runtime_preprocessors: ( Sequence[preprocessors_lib.PyGrainAirIOPreprocessor] | None ) = None, @@ -581,7 +581,7 @@ def build(self) -> GrainTask: ) @classmethod - def from_task(cls, task: GrainTask) -> "GrainTaskBuilder": + def from_task(cls, task: GrainTask) -> "GrainTaskBuilder": # pyrefly: ignore[bad-override] """Returns TaskBuilder for the given existing Task object. This method takes an existing task, copies its properties into a new diff --git a/airio/_src/pygrain/preprocessors.py b/airio/_src/pygrain/preprocessors.py index 321db93..99d0429 100644 --- a/airio/_src/pygrain/preprocessors.py +++ b/airio/_src/pygrain/preprocessors.py @@ -42,7 +42,7 @@ class MapFnTransform( def map(self, element): """Maps a single element.""" return core_preprocessors.inject_runtime_args_to_fn( - self.map_fn, self.runtime_args + self.map_fn, self.runtime_args # pyrefly: ignore[bad-argument-type] )(element) @@ -57,7 +57,7 @@ def random_map(self, element, rng: np.random.Generator): """Maps a single element.""" jax_rng = jax.random.key(rng.integers(0, 2**16 - 1)) return core_preprocessors.inject_runtime_args_to_fn( - self.map_fn, self.runtime_args + self.map_fn, self.runtime_args # pyrefly: ignore[bad-argument-type] )(element, jax_rng) @@ -71,7 +71,7 @@ class FilterFnTransform( def filter(self, element) -> bool: """Filters a single element.""" return core_preprocessors.inject_runtime_args_to_fn( - self.filter_fn, self.runtime_args + self.filter_fn, self.runtime_args # pyrefly: ignore[bad-argument-type] )(element) diff --git a/airio/_src/pygrain/tokenizer.py b/airio/_src/pygrain/tokenizer.py index 51a916c..7dd98da 100644 --- a/airio/_src/pygrain/tokenizer.py +++ b/airio/_src/pygrain/tokenizer.py @@ -44,7 +44,7 @@ def __call__(self, orig_example: Inp) -> Out: pad_width = [(0, 1)] # Tokenized rank is generally 1; adjust pad_width in case it's more pad_width += [(0, 0)] * (len(encoded_val.shape) - 1) - encoded_val = np.pad( + encoded_val = np.pad( # pyrefly: ignore[no-matching-overload] encoded_val, pad_width, constant_values=tokenizer_config.vocab.eos_id, diff --git a/airio/examples/feature_converter.py b/airio/examples/feature_converter.py index 99fd96a..09d6dd4 100644 --- a/airio/examples/feature_converter.py +++ b/airio/examples/feature_converter.py @@ -47,10 +47,10 @@ def _imdb_preprocessor(raw_example: Dict[str, bytes]) -> Dict[str, str]: tfds_name="imdb_reviews/plain_text:1.0.0", splits=["train"] ), preprocessors=[ - airio.MapFnTransform(_imdb_preprocessor), + airio.MapFnTransform(_imdb_preprocessor), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs={ + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs={ # pyrefly: ignore[unexpected-keyword] "inputs": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), "targets": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), }, diff --git a/airio/examples/inspect.py b/airio/examples/inspect.py index 9d6f9f1..4e48b56 100644 --- a/airio/examples/inspect.py +++ b/airio/examples/inspect.py @@ -47,10 +47,10 @@ def _imdb_preprocessor(raw_example: Dict[str, bytes]) -> Dict[str, str]: tfds_name="imdb_reviews/plain_text:1.0.0", splits=["train"] ), preprocessors=[ - airio.MapFnTransform(_imdb_preprocessor), + airio.MapFnTransform(_imdb_preprocessor), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs={ + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs={ # pyrefly: ignore[unexpected-keyword] "inputs": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), "targets": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), }, diff --git a/airio/examples/mixtures.py b/airio/examples/mixtures.py index 74dbd3e..7c3a5c8 100644 --- a/airio/examples/mixtures.py +++ b/airio/examples/mixtures.py @@ -75,11 +75,11 @@ def get_mc4_mixture( tfds_name="c4/multilingual:3.1.0", splits={"train": lang, "validation": f"{lang}-validation"}, ), - preprocessors=[ - airio.MapFnTransform(rekey_fn), + preprocessors=[ # pyrefly: ignore[bad-argument-type] + airio.MapFnTransform(rekey_fn), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs=tokenizer_configs, + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs=tokenizer_configs, # pyrefly: ignore[unexpected-keyword] ) ), airio_common.span_corruption.create_span_corruption_transform( diff --git a/airio/examples/quickstart.py b/airio/examples/quickstart.py index 0d33aa1..d4b36b6 100644 --- a/airio/examples/quickstart.py +++ b/airio/examples/quickstart.py @@ -46,10 +46,10 @@ def _imdb_preprocessor(raw_example: Dict[str, bytes]) -> Dict[str, str]: tfds_name="imdb_reviews/plain_text:1.0.0", splits=["train"] ), preprocessors=[ - airio.MapFnTransform(_imdb_preprocessor), + airio.MapFnTransform(_imdb_preprocessor), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs={ + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs={ # pyrefly: ignore[unexpected-keyword] "inputs": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), "targets": airio.TokenizerConfig(vocab=DEFAULT_VOCAB), }, diff --git a/airio/examples/tasks.py b/airio/examples/tasks.py index cbcdb95..f3690e8 100644 --- a/airio/examples/tasks.py +++ b/airio/examples/tasks.py @@ -54,16 +54,16 @@ def get_wmt_19_ende_v003_task( ), preprocessors=[ airio.MapFnTransform( - functools.partial( + functools.partial( # pyrefly: ignore[bad-argument-count] translate, source_language=builder_config.language_pair[1], target_language=builder_config.language_pair[0], ) ), airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs=tokenizer_configs, - copy_pretokenized=False, + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs=tokenizer_configs, # pyrefly: ignore[unexpected-keyword] + copy_pretokenized=False, # pyrefly: ignore[unexpected-keyword] ) ), ], @@ -92,11 +92,11 @@ def get_nqo_v001_task( tfds_name=tfds_name, splits=["train", "validation"] ), preprocessors=[ - airio.MapFnTransform(question), + airio.MapFnTransform(question), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs=tokenizer_configs, - copy_pretokenized=False, + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs=tokenizer_configs, # pyrefly: ignore[unexpected-keyword] + copy_pretokenized=False, # pyrefly: ignore[unexpected-keyword] ) ), ], @@ -140,12 +140,14 @@ def translate( } if isinstance(ex[source_language], bytes): src_str = ( + # pyrefly: ignore[unsupported-operation] f"translate {lang_id_to_string[source_language]} to" f" {lang_id_to_string[target_language]}: ".encode() + ex[source_language] ) else: src_str = ( + # pyrefly: ignore[unsupported-operation] f"translate {lang_id_to_string[source_language]} to" f" {lang_id_to_string[target_language]}: " + ex[source_language] @@ -214,16 +216,16 @@ def get_c4_v220_span_corruption_task( source=airio.TfdsDataSource( tfds_name="c4/en:2.2.0", splits=["train", "validation"] ), - preprocessors=[ - airio.MapFnTransform(rekey_fn), + preprocessors=[ # pyrefly: ignore[bad-argument-type] + airio.MapFnTransform(rekey_fn), # pyrefly: ignore[bad-argument-count] airio.MapFnTransform( - airio.Tokenizer( - tokenizer_configs=tokenizer_configs, + airio.Tokenizer( # pyrefly: ignore[bad-argument-count] + tokenizer_configs=tokenizer_configs, # pyrefly: ignore[unexpected-keyword] ) ), airio_common.span_corruption.create_span_corruption_transform( tokenizer_configs ), - airio.MapFnTransform(append_eos_after_trim_fn), + airio.MapFnTransform(append_eos_after_trim_fn), # pyrefly: ignore[bad-argument-count] ], ) diff --git a/airio/examples/train_wmt.py b/airio/examples/train_wmt.py index b87d95f..3d2daad 100644 --- a/airio/examples/train_wmt.py +++ b/airio/examples/train_wmt.py @@ -62,8 +62,8 @@ def get_t5_model(**config_overrides) -> models.EncoderDecoderModel: tiny_config = dataclasses.replace(tiny_config, **config_overrides) return models.EncoderDecoderModel( # pytype: disable=wrong-arg-types module=network.Transformer(tiny_config), - input_vocabulary=_DEFAULT_VOCAB, - output_vocabulary=_DEFAULT_VOCAB, + input_vocabulary=_DEFAULT_VOCAB, # pyrefly: ignore[bad-argument-type] + output_vocabulary=_DEFAULT_VOCAB, # pyrefly: ignore[bad-argument-type] optimizer_def=adafactor.Adafactor( decay_rate=0.8, step_offset=0,