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
10 changes: 5 additions & 5 deletions tensorflow_datasets/core/as_dataframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,11 +111,11 @@ def _get_feature(
if type(feature) != features.Sequence and not path: # pylint: disable=unidiomatic-typecheck
break
sequence_rank += 1
feature = feature.feature # Extract inner feature # pytype: disable=attribute-error
feature = feature.feature # Extract inner feature

if path: # Has level deeper, recurse
feature = typing.cast(features.FeaturesDict, feature)
feature, nested_sequence_rank = _get_feature(path[1:], feature[path[0]]) # pytype: disable=wrong-arg-types
feature, nested_sequence_rank = _get_feature(path[1:], feature[path[0]])
sequence_rank += nested_sequence_rank

return feature, sequence_rank
Expand Down Expand Up @@ -186,7 +186,7 @@ class StyledDataFrame(pd.DataFrame):
# selecting sub-data frames.

def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs) # pytype: disable=wrong-arg-count # re-none
super().__init__(*args, **kwargs)
# Use name-mangling for forward-compatibility in case pandas
# adds a `_styler` attribute in the future.
self.__styler: Optional[Styler] = None
Expand All @@ -195,13 +195,13 @@ def __init__(self, *args, **kwargs):
def current_style(self) -> Styler:
"""Like `pandas.DataFrame.style`, but attach the style to the DataFrame."""
if self.__styler is None:
self.__styler = super().style # pytype: disable=attribute-error # re-none
self.__styler = super().style
return self.__styler

def _repr_html_(self) -> Union[None, str]:
# See base class for doc
if self.__styler is None:
return super()._repr_html_() # pytype: disable=attribute-error # re-none
return super()._repr_html_()
return self.__styler._repr_html_() # pylint: disable=protected-access

# Pack `as_supervised=True` datasets
Expand Down
6 changes: 3 additions & 3 deletions tensorflow_datasets/core/dataset_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -1478,7 +1478,7 @@ def _get_filename_template(
split=split_name,
dataset_name=self.name,
data_dir=self.data_path,
filetype_suffix=self.info.file_format.file_suffix, # pytype: disable=attribute-error
filetype_suffix=self.info.file_format.file_suffix,
)


Expand Down Expand Up @@ -1806,7 +1806,7 @@ def _generate_splits(
# Finalize the splits (after apache beam completed, if it was used)
return [future.result() for future in split_info_futures]

def _download_and_prepare( # pytype: disable=signature-mismatch # overriding-parameter-type-checks
def _download_and_prepare( # pyrefly: ignore[bad-override]
self,
dl_manager: download.DownloadManager,
download_config: download.DownloadConfig,
Expand Down Expand Up @@ -1931,7 +1931,7 @@ def _download_and_prepare(
) -> None:
download_config = download_config or download.DownloadConfig()

split_builder = split_builder_lib.SplitBuilder( # pytype: disable=wrong-arg-types
split_builder = split_builder_lib.SplitBuilder(
split_dict=self.info.splits,
features=self.info.features,
dataset_size=self.info.dataset_size,
Expand Down
6 changes: 3 additions & 3 deletions tensorflow_datasets/core/dataset_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -420,7 +420,7 @@ def set_nondeterministic_order(self, nondeterministic_order: bool) -> None:
def homepage(self) -> str:
urls = self.as_proto.location.urls
tfds_homepage = f"https://www.tensorflow.org/datasets/catalog/{self.name}"
return urls and urls[0] or tfds_homepage # pytype: disable=bad-return-type
return urls and urls[0] or tfds_homepage # pyrefly: ignore[bad-return]

@property
def citation(self) -> str:
Expand Down Expand Up @@ -726,7 +726,7 @@ def read_from_directory(self, dataset_info_dir: epath.PathLike) -> None:
)

# Update splits
filename_template = naming.ShardedFileTemplate( # pytype: disable=wrong-arg-types # always-use-property-annotation
filename_template = naming.ShardedFileTemplate(
dataset_name=self.name,
data_dir=self.data_dir, # pyrefly: ignore[bad-argument-type]
filetype_suffix=parsed_proto.file_format or "tfrecord",
Expand Down Expand Up @@ -1202,7 +1202,7 @@ def pack_as_supervised_ds(
and isinstance(ds.element_spec, tuple)
and len(ds.element_spec) == 2
):
x_key, y_key = ds_info.supervised_keys # pytype: disable=bad-unpacking
x_key, y_key = ds_info.supervised_keys # pyrefly: ignore[bad-unpacking]
ds = ds.map(lambda x, y: {x_key: x, y_key: y})
return ds
else: # If dataset isn't a supervised tuple (input, label), return as-is
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/dataset_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,7 @@ def _elem_to_numpy_eager(
) -> Union[NumpyElem, Iterable[NumpyElem]]:
"""Converts a single element from tf to numpy."""
if isinstance(tf_el, tf.Tensor):
return tf_el._numpy() # pytype: disable=attribute-error # pylint: disable=protected-access
return tf_el._numpy() # pylint: disable=protected-access
elif isinstance(tf_el, tf.RaggedTensor):
return tf_el
elif isinstance(tf_el, tf.data.Dataset):
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/example_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ def __post_init__(self):
def parse_example(
self, serialized_example: bytes | memoryview
) -> Mapping[str, Union[np.ndarray, list[Any]]]:
example = tf_example_pb2.Example.FromString(serialized_example) # pyrefly: ignore[bad-argument-type]
example = tf_example_pb2.Example.FromString(serialized_example)
np_example = _features_to_numpy(example.features, self._flat_example_specs) # pyrefly: ignore[bad-argument-type]
return utils.pack_as_nest_dict(np_example, self.example_specs)

Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/features/audio_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
lazy_decode=value.lazy_decode or False,
)

def to_json_content(self) -> feature_pb2.AudioFeature: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> feature_pb2.AudioFeature: # pyrefly: ignore[bad-override]
return feature_pb2.AudioFeature(
shape=feature_lib.to_shape_proto(self.shape),
dtype=feature_lib.dtype_to_str(self.dtype), # pyrefly: ignore[bad-argument-type]
Expand Down
4 changes: 1 addition & 3 deletions tensorflow_datasets/core/features/bounding_boxes.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,9 +152,7 @@ def from_json_content( # pyrefly: ignore[bad-override]

def to_json_content( # pyrefly: ignore[bad-override]
self,
) -> (
feature_pb2.BoundingBoxFeature
): # pytype: disable=signature-mismatch # overriding-return-type-checks
) -> feature_pb2.BoundingBoxFeature:
bbox_format = None
if self.bbox_format:
bbox_format = (
Expand Down
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/features/class_label_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,7 +190,7 @@ def load_metadata(self, data_dir, feature_name=None) -> Optional[list[str]]:
pass

def _additional_repr_info(self) -> dict[str, int]:
return {"num_classes": self.num_classes} # pytype: disable=bad-return-type # always-use-property-annotation
return {"num_classes": self.num_classes} # pyrefly: ignore[bad-assignment]

def repr_html(self, ex: int) -> str: # pyrefly: ignore[bad-override]
"""Class labels are displayed with their name."""
Expand All @@ -209,7 +209,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
return cls(**value) # pyrefly: ignore[bad-argument-type]
return cls(num_classes=value.num_classes)

def to_json_content(self) -> feature_pb2.ClassLabel: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> feature_pb2.ClassLabel: # pyrefly: ignore[bad-override]
return feature_pb2.ClassLabel(num_classes=self.num_classes)

@classmethod
Expand Down
6 changes: 3 additions & 3 deletions tensorflow_datasets/core/features/feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -409,7 +409,7 @@ def cls_from_name(cls, python_class_name: str) -> Type['FeatureConnector']:
raise ValueError(
f'Python class name must contain a dot, got: "{python_class_name}"'
)
module_name, _ = python_class_name.rsplit('.', maxsplit=1) # pytype: disable=attribute-error
module_name, _ = python_class_name.rsplit('.', maxsplit=1)
try:
# Import to register the FeatureConnector
importlib.import_module(module_name)
Expand Down Expand Up @@ -570,7 +570,7 @@ def from_json_content(
"""
if not isinstance(value, dict):
raise TypeError(f'Unexpected feature connector value: {value!r}')
return cls(doc=doc, **value) # pytype: disable=not-instantiable
return cls(doc=doc, **value)

def to_json_content(self) -> Union[Json, message.Message]:
"""FeatureConnector factory (to overwrite).
Expand Down Expand Up @@ -1104,7 +1104,7 @@ def _has_shape_ambiguity(in_shape: Shape, out_shape: Shape) -> bool:
"""Returns True if the shape can be an empty sequence with unknown shape."""
# Normalize shape if running with `tf.compat.v1.disable_v2_tensorshape`
if isinstance(in_shape, tf.TensorShape):
in_shape = in_shape.as_list() # pytype: disable=attribute-error
in_shape = in_shape.as_list()

return bool(
in_shape[0] is None # Empty sequence
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/features/image_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -417,7 +417,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
) -> 'Image':
if isinstance(value, dict):
# For backwards compatibility
return cls( # pytype: disable=wrong-arg-types
return cls(
shape=tuple(value['shape']), # pyrefly: ignore[bad-argument-type]
dtype=feature_lib.dtype_from_str(value['dtype']), # pyrefly: ignore[bad-argument-type]
encoding_format=value['encoding_format'], # pyrefly: ignore[bad-argument-type]
Expand Down
6 changes: 3 additions & 3 deletions tensorflow_datasets/core/features/sequence_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,10 +174,10 @@ def load_metadata(self, *args, **kwargs):

def __getitem__(self, key):
"""Convenience method to access the underlying features."""
return self._feature[key] # pytype: disable=unsupported-operands
return self._feature[key]

def __contains__(self, key: str) -> bool:
return key in self._feature # pytype: disable=unsupported-operands
return key in self._feature

def __getattr__(self, key):
"""Allow to access the underlying attributes directly."""
Expand Down Expand Up @@ -318,5 +318,5 @@ def update_length(elem):
# 3. Extract each individual elements
return [
utils.map_nested(lambda elem: elem[i], dict_list, dict_only=True) # pylint: disable=cell-var-from-loop
for i in range(length['value']) # pytype: disable=wrong-arg-types
for i in range(length['value']) # pyrefly: ignore[bad-argument-type]
]
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/features/text_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,7 @@ def save_metadata(self, data_dir, feature_name: str) -> None: # pyrefly: ignore
def load_metadata(self, data_dir, feature_name: str) -> None: # pyrefly: ignore[bad-override]
if self._encoder_cls:
fname_prefix = _file_name_prefix_for_metadata(feature_name, data_dir)
self._encoder = self._encoder_cls.load_from_file(fname_prefix) # pytype: disable=attribute-error
self._encoder = self._encoder_cls.load_from_file(fname_prefix)
return

# Error checking: ensure there are no metadata files
Expand Down Expand Up @@ -209,7 +209,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
)
return cls()

def to_json_content(self) -> Union[Json, feature_pb2.TextFeature]: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> Union[Json, feature_pb2.TextFeature]: # pyrefly: ignore[bad-override]
if self._encoder:
logging.warning(
"Dataset is using deprecated text encoder API which will be removed "
Expand Down
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/features/translation_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
assert not value.variable_languages_per_example
return cls(languages=value.languages)

def to_json_content(self) -> feature_pb2.TranslationFeature: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> feature_pb2.TranslationFeature: # pyrefly: ignore[bad-override]
if self._encoder or self._encoder_config:
raise ValueError(
"TFDS encoder are deprecated and will be removed soon. "
Expand Down Expand Up @@ -249,7 +249,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
assert value.variable_languages_per_example
return cls(languages=value.languages)

def to_json_content(self) -> feature_pb2.TranslationFeature: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> feature_pb2.TranslationFeature: # pyrefly: ignore[bad-override]
return feature_pb2.TranslationFeature(
languages=self.languages, variable_languages_per_example=True
)
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/features/video_feature.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,7 +207,7 @@ def from_json_content( # pyrefly: ignore[bad-override]
ffmpeg_extra_args=value.ffmpeg_extra_args,
)

def to_json_content(self) -> feature_pb2.VideoFeature: # pytype: disable=signature-mismatch # overriding-return-type-checks
def to_json_content(self) -> feature_pb2.VideoFeature: # pyrefly: ignore[bad-override]
return feature_pb2.VideoFeature(
shape=feature_lib.to_shape_proto(self.shape),
dtype=feature_lib.dtype_to_str(self.dtype), # pyrefly: ignore[bad-argument-type]
Expand All @@ -219,5 +219,5 @@ def to_json_content(self) -> feature_pb2.VideoFeature: # pytype: disable=signat
def repr_html(self, ex: np.ndarray) -> str:
"""Video are displayed as `<video>`."""
return image_feature.make_video_repr_html(
ex, use_colormap=self.feature._use_colormap # pylint: disable=protected-access # pytype: disable=attribute-error
ex, use_colormap=self.feature._use_colormap # pylint: disable=protected-access
)
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/file_adapters.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,7 +283,7 @@ def make_tf_data(
buffer_size: int | None = None,
) -> tf.data.Dataset:
buffer_size = buffer_size or cls.BUFFER_SIZE
from riegeli.tensorflow.ops import riegeli_dataset_ops as riegeli_tf # pylint: disable=g-import-not-at-top # pytype: disable=import-error
from riegeli.tensorflow.ops import riegeli_dataset_ops as riegeli_tf # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

return riegeli_tf.RiegeliDataset(filename, buffer_size=buffer_size)

Expand Down
6 changes: 3 additions & 3 deletions tensorflow_datasets/core/load.py
Original file line number Diff line number Diff line change
Expand Up @@ -242,7 +242,7 @@ def get_dataset_repr() -> str:
with py_utils.try_reraise(
prefix=f'Failed to construct {get_dataset_repr()}: '
):
return cls(**builder_kwargs) # pytype: disable=not-instantiable
return cls(**builder_kwargs)

# If neither the code nor the files are found, raise DatasetNotFoundError
if not_found_error is not None:
Expand Down Expand Up @@ -410,7 +410,7 @@ def load_dataset(
raise RuntimeError(
f'Unsupported return type {type(load_output)} of `load` function.'
)
return loaded_datasets # pytype: disable=bad-return-type
return loaded_datasets

def load_datasets(
self,
Expand Down Expand Up @@ -941,7 +941,7 @@ def single_full_names(
_iter_single_full_names(
builder_name,
builder_cls(builder_name),
current_version_only=current_version_only, # pytype: disable=wrong-arg-types
current_version_only=current_version_only,
)
)

Expand Down
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/registered.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,7 +226,7 @@ def imported_dataset_collection_cls(

dataset_collection_cls = _DATASET_COLLECTION_REGISTRY[name]

return dataset_collection_cls # pytype: disable=bad-return-type
return dataset_collection_cls


class RegisteredDataset(abc.ABC):
Expand Down Expand Up @@ -478,7 +478,7 @@ def imported_builder_cls(name: str) -> Type[RegisteredDataset]:
if name in _ABSTRACT_DATASET_REGISTRY:
# Will raise TypeError: Can't instantiate abstract class X with abstract
# methods y, before __init__ even get called
_ABSTRACT_DATASET_REGISTRY[name]() # pytype: disable=not-callable
_ABSTRACT_DATASET_REGISTRY[name]()
# Alternatively, could manually extract the list of non-implemented
# abstract methods.
raise AssertionError(f'Dataset {name} is an abstract class.')
Expand Down
4 changes: 2 additions & 2 deletions tensorflow_datasets/core/shuffle.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,7 +288,7 @@ def add(self, key: type_utils.Key, data: bytes) -> bool: # pyrefly: ignore[bad-
hkey = self._hasher.hash_key(key)
if self._ignore_duplicates:
if hkey in self._seen_keys:
return # pytype: disable=bad-return-type
return # pyrefly: ignore[bad-return]
self._seen_keys.add(hkey)
if self._disable_shuffling:
# Use the original key and not the hashed key to maintain the order.
Expand All @@ -298,7 +298,7 @@ def add(self, key: type_utils.Key, data: bytes) -> bool: # pyrefly: ignore[bad-
self._add_to_mem_buffer(hkey, data)
else:
self._add_to_bucket(hkey, data)
self._num_examples += 1 # pytype: disable=bad-return-type
self._num_examples += 1

def __iter__(self) -> Iterator[type_utils.KeySerializedExample]:
self._read_only = True
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/core/visibility.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ def _set_default_visibility() -> None:
If the script executed is a TFDS script, then it restricts the visibility
to only open-source non-community datasets.
"""
import __main__ # pytype: disable=import-error # pylint: disable=g-import-not-at-top
import __main__ # pylint: disable=g-import-not-at-top

main_file = getattr(__main__, '__file__', None)
if main_file and 'tensorflow_datasets' in pathlib.Path(main_file).parts:
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/text/c4.py
Original file line number Diff line number Diff line change
Expand Up @@ -554,7 +554,7 @@ def _split_generators(
_OPENWEBTEXT_URLS_ZIP,
)
)
file_paths["openwebtext_urls_zip"] = dl_manager.extract(owt_path) # pyrefly: ignore[unsupported-operation]
file_paths["openwebtext_urls_zip"] = dl_manager.extract(owt_path)

file_paths = tree.map_structure(os.fspath, file_paths)

Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/text/c4_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,7 +400,7 @@ def _validate_features(page):
fileobj=f
) as g:
page = PageFeatures()
for i, line in enumerate(io.TextIOWrapper(g, encoding="utf-8")): # pytype: disable=wrong-arg-types
for i, line in enumerate(io.TextIOWrapper(g, encoding="utf-8")):
line = line.strip()
if not line:
continue
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/text/glue.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ def __init__(
"2.0.0": "Update data source for glue/qqp.",
},
**kwargs,
) # pytype: disable=wrong-arg-types # gen-stub-imports
)
self.text_features = text_features
self.label_column = label_column
self.label_classes = label_classes
Expand Down
4 changes: 2 additions & 2 deletions tensorflow_datasets/text/super_glue.py
Original file line number Diff line number Diff line change
Expand Up @@ -636,10 +636,10 @@ def _fix_span_text(k):
return

if "theyscold" in text:
ex["text"].replace("theyscold", "they scold") # pytype: disable=attribute-error
ex["text"].replace("theyscold", "they scold")
ex["span2_index"] = 10
# Make sure case of the first words match.
first_word = ex["text"].split()[index] # pytype: disable=attribute-error
first_word = ex["text"].split()[index]
if first_word[0].islower():
text = text[0].lower() + text[1:]
else:
Expand Down
2 changes: 1 addition & 1 deletion tensorflow_datasets/text/wikipedia.py
Original file line number Diff line number Diff line change
Expand Up @@ -531,7 +531,7 @@ def _parse_and_clean_wikicode(raw_content):
wikicode = mwparserfromhell.parse(raw_content)

def rm_wikilink(obj):
return bool(re_rm_wikilink.match(str(obj.title))) # pytype: disable=wrong-arg-types
return bool(re_rm_wikilink.match(str(obj.title)))

def rm_tag(obj):
return str(obj.tag) in {"ref", "table"}
Expand Down
Loading