A `Dataset` that randomly shuffles the elements of its input.
| 3150 | |
| 3151 | |
| 3152 | class ShuffleDataset(UnaryUnchangedStructureDataset): |
| 3153 | """A `Dataset` that randomly shuffles the elements of its input.""" |
| 3154 | |
| 3155 | def __init__(self, |
| 3156 | input_dataset, |
| 3157 | buffer_size, |
| 3158 | seed=None, |
| 3159 | reshuffle_each_iteration=None): |
| 3160 | """Randomly shuffles the elements of this dataset. |
| 3161 | |
| 3162 | Args: |
| 3163 | input_dataset: The input dataset. |
| 3164 | buffer_size: A `tf.int64` scalar `tf.Tensor`, representing the number of |
| 3165 | elements from this dataset from which the new dataset will sample. |
| 3166 | seed: (Optional.) A `tf.int64` scalar `tf.Tensor`, representing the random |
| 3167 | seed that will be used to create the distribution. See |
| 3168 | `tf.compat.v1.set_random_seed` for behavior. |
| 3169 | reshuffle_each_iteration: (Optional.) A boolean, which if true indicates |
| 3170 | that the dataset should be pseudorandomly reshuffled each time it is |
| 3171 | iterated over. (Defaults to `True`.) |
| 3172 | |
| 3173 | Returns: |
| 3174 | A `Dataset`. |
| 3175 | |
| 3176 | Raises: |
| 3177 | ValueError: if invalid arguments are provided. |
| 3178 | """ |
| 3179 | self._input_dataset = input_dataset |
| 3180 | self._buffer_size = ops.convert_to_tensor( |
| 3181 | buffer_size, dtype=dtypes.int64, name="buffer_size") |
| 3182 | self._seed, self._seed2 = random_seed.get_seed(seed) |
| 3183 | |
| 3184 | if reshuffle_each_iteration is None: |
| 3185 | self._reshuffle_each_iteration = True |
| 3186 | else: |
| 3187 | self._reshuffle_each_iteration = reshuffle_each_iteration |
| 3188 | |
| 3189 | if tf2.enabled() and self._reshuffle_each_iteration and ( |
| 3190 | context.executing_eagerly() or |
| 3191 | ops.get_default_graph()._building_function): # pylint: disable=protected-access |
| 3192 | self._seed_generator = _RandomSeedGenerator(self._seed, self._seed2) |
| 3193 | variant_tensor = gen_dataset_ops.shuffle_dataset_v2( |
| 3194 | input_dataset._variant_tensor, # pylint: disable=protected-access |
| 3195 | buffer_size=self._buffer_size, |
| 3196 | seed_generator=self._seed_generator.handle, |
| 3197 | **self._flat_structure) |
| 3198 | else: |
| 3199 | variant_tensor = gen_dataset_ops.shuffle_dataset( |
| 3200 | input_dataset._variant_tensor, # pylint: disable=protected-access |
| 3201 | buffer_size=self._buffer_size, |
| 3202 | seed=self._seed, |
| 3203 | seed2=self._seed2, |
| 3204 | reshuffle_each_iteration=self._reshuffle_each_iteration, |
| 3205 | **self._flat_structure) |
| 3206 | super(ShuffleDataset, self).__init__(input_dataset, variant_tensor) |
| 3207 | |
| 3208 | |
| 3209 | class TakeDataset(UnaryUnchangedStructureDataset): |