MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / ShuffleDataset

Class ShuffleDataset

tensorflow/python/data/ops/dataset_ops.py:3152–3206  ·  view source on GitHub ↗

A `Dataset` that randomly shuffles the elements of its input.

Source from the content-addressed store, hash-verified

3150
3151
3152class 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
3209class TakeDataset(UnaryUnchangedStructureDataset):

Callers 1

shuffleMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected