diff --git a/mediapipe/model_maker/python/core/data/dataset.py b/mediapipe/model_maker/python/core/data/dataset.py index a92b05c0..11396938 100644 --- a/mediapipe/model_maker/python/core/data/dataset.py +++ b/mediapipe/model_maker/python/core/data/dataset.py @@ -18,7 +18,7 @@ from __future__ import division from __future__ import print_function import functools -from typing import Callable, Optional, Tuple, TypeVar +from typing import Any, Callable, Optional, Tuple, TypeVar # Dependency imports import tensorflow as tf @@ -66,12 +66,14 @@ class Dataset(object): """ return self._size - def gen_tf_dataset(self, - batch_size: int = 1, - is_training: bool = False, - shuffle: bool = False, - preprocess: Optional[Callable[..., bool]] = None, - drop_remainder: bool = False) -> tf.data.Dataset: + def gen_tf_dataset( + self, + batch_size: int = 1, + is_training: bool = False, + shuffle: bool = False, + preprocess: Optional[Callable[..., Any]] = None, + drop_remainder: bool = False, + ) -> tf.data.Dataset: """Generates a batched tf.data.Dataset for training/evaluation. Args: diff --git a/mediapipe/model_maker/python/core/tasks/classifier.py b/mediapipe/model_maker/python/core/tasks/classifier.py index abcfff83..bfe0f027 100644 --- a/mediapipe/model_maker/python/core/tasks/classifier.py +++ b/mediapipe/model_maker/python/core/tasks/classifier.py @@ -48,11 +48,13 @@ class Classifier(custom_model.CustomModel): self._hparams: hp.BaseHParams = None self._history: tf.keras.callbacks.History = None - def _train_model(self, - train_data: classification_ds.ClassificationDataset, - validation_data: classification_ds.ClassificationDataset, - preprocessor: Optional[Callable[..., bool]] = None, - checkpoint_path: Optional[str] = None): + def _train_model( + self, + train_data: classification_ds.ClassificationDataset, + validation_data: classification_ds.ClassificationDataset, + preprocessor: Optional[Callable[..., Any]] = None, + checkpoint_path: Optional[str] = None, + ): """Trains the classifier model. Compiles and fits the tf.keras `_model` and records the `_history`. diff --git a/mediapipe/model_maker/python/core/utils/model_util.py b/mediapipe/model_maker/python/core/utils/model_util.py index 69a8654e..7a0b8fcf 100644 --- a/mediapipe/model_maker/python/core/utils/model_util.py +++ b/mediapipe/model_maker/python/core/utils/model_util.py @@ -115,9 +115,11 @@ def get_steps_per_epoch(steps_per_epoch: Optional[int] = None, def convert_to_tflite( model: tf.keras.Model, quantization_config: Optional[quantization.QuantizationConfig] = None, - supported_ops: Tuple[tf.lite.OpsSet, - ...] = (tf.lite.OpsSet.TFLITE_BUILTINS,), - preprocess: Optional[Callable[..., bool]] = None) -> bytearray: + supported_ops: Tuple[tf.lite.OpsSet, ...] = ( + tf.lite.OpsSet.TFLITE_BUILTINS, + ), + preprocess: Optional[Callable[..., Any]] = None, +) -> bytearray: """Converts the input Keras model to TFLite format. Args: