diff --git a/recml/core/data/tf_dataset_factory.py b/recml/core/data/tf_dataset_factory.py index c849022..f7e0ab4 100644 --- a/recml/core/data/tf_dataset_factory.py +++ b/recml/core/data/tf_dataset_factory.py @@ -209,6 +209,17 @@ class TFDatasetFactory(types.Factory[tf.data.Dataset]): postprocessors: A sequence of postprocessing functions to apply to the dataset. These will be applied in the order they are provided after the preprocessor if specified. + rebatch_after_transform: Whether to re-batch the dataset to the per-process + batch size after applying the transformations (`tf_transform_output`, + `preprocessor`, and `postprocessors`). This is useful when the + transformations change the number of examples in a batch, e.g. by + filtering or expanding examples. If True, the batch size is not enforced + via `tf.ensure_shape` inside the transform function; instead the dataset + is re-batched with `drop_remainder` after the transformations. Note that + the leading (batch) dimension of the transformed outputs may differ from + the original batch size and vary across batches, but within a single + batch all output tensors must share the same leading dimension so that + they can be split into individual examples. Defaults to False. label_name: The name of the label feature. If passed, this will be popped from the features dictionary after performing any transformations and the dataset returned will consist of tuples of the features dictionary and the @@ -257,6 +268,7 @@ class TFDatasetFactory(types.Factory[tf.data.Dataset]): filter_fn: FilterFn | None = None preprocessor: FeatureTransformationFn | None = None postprocessors: Sequence[FeatureTransformationFn] = () + rebatch_after_transform: bool = False label_name: str | None = None data_options: tf.data.Options | None = None sharding_info: DatasetShardingInfo = dataclasses.field( @@ -447,6 +459,14 @@ def map_fns(self) -> Sequence[TransformFn]: and self.label_name is None ): return [] + # When re-batching after the transformations, the transformations may + # change the batch dimension, so don't enforce the batch size inside the + # transform function. + batch_size = None + if not self.rebatch_after_transform: + batch_size = self.sharding_info.per_process_batch_size( + self.global_batch_size + ) return [ build_transform_fn( tf_transform_output=self.tf_transform_output, @@ -455,9 +475,7 @@ def map_fns(self) -> Sequence[TransformFn]: ) + list(self.postprocessors), label_name=self.label_name, - batch_size=self.sharding_info.per_process_batch_size( - self.global_batch_size - ), + batch_size=batch_size, ) ] @@ -558,6 +576,23 @@ def _maybe_filter_dataset(self, dataset: tf.data.Dataset) -> tf.data.Dataset: ) return dataset + def _maybe_rebatch_after_transform( + self, dataset: tf.data.Dataset + ) -> tf.data.Dataset: + """Re-batches a transformed dataset to the per-process batch size.""" + if not self.rebatch_after_transform or not self.map_fns: + return dataset + per_process_batch_size = self.sharding_info.per_process_batch_size( + self.global_batch_size + ) + logging.info( + "Re-batching the dataset after transformations to batch size: %s", + per_process_batch_size, + ) + return dataset.rebatch( + per_process_batch_size, drop_remainder=self.drop_remainder + ) + def _maybe_shuffle_and_repeat(self, dataset: tf.data.Dataset): """Shuffles and / or repeats an examples dataset.""" if self.shuffle: @@ -611,6 +646,9 @@ def make(self) -> tf.data.Dataset: for fn in self.map_fns: dataset = dataset.map(fn, num_parallel_calls=self.num_parallel_threads) + # Re-batch the dataset in case the transformations changed the batch size. + dataset = self._maybe_rebatch_after_transform(dataset) + # Apply TF Data Service after preprocessing. if self.offload_preprocessing_to_tf_data_service: dataset = self._maybe_apply_tf_data_service(dataset)