Skip to content
Merged
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
44 changes: 41 additions & 3 deletions recml/core/data/tf_dataset_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand All @@ -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,
)
]

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading