Метод rebatch
Метод rebatch применяется к датасету,
элементы которого уже сгруппированы в батчи,
и создает новый датасет с батчами другого размера.
Первым параметром метод принимает желаемый размер
батча batch_size. Вторым необязательным
параметром можно передать drop_remainder -
логическое значение, которое указывает, отбрасывать
ли последний неполный батч. В отличие от метода
batch, который группирует отдельные элементы,
метод rebatch работает поверх уже
существующих батчей и не разбивает их на отдельные
элементы.
Синтаксис
Dataset.rebatch(batch_size, [drop_remainder])
Пример
Давайте создадим датасет из тензора, сгруппируем
элементы в батчи по 2, а затем перегруппируем
их в батчи по 3:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices(tf.constant([1, 2, 3, 4, 5, 6, 7, 8]))
ds = ds.batch(2)
ds = ds.rebatch(3)
for batch in ds:
print(batch)
Результат выполнения кода:
tf.Tensor([1 2 3], shape=(3,), dtype=int32)
tf.Tensor([4 5 6], shape=(3,), dtype=int32)
tf.Tensor([7 8], shape=(2,), dtype=int32)
Пример
Давайте перегруппируем батчи с параметром
drop_remainder, равным True,
чтобы отбросить последний неполный батч:
Результат выполнения кода:
tf.Tensor([1 2 3], shape=(3,), dtype=int32)
tf.Tensor([4 5 6], shape=(3,), dtype=int32)
Пример
Давайте создадим датасет из двумерного тензора,
сгруппируем его в батчи по 2 и
перегруппируем в батчи по 3:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12]])
ds = tf.data.Dataset.from_tensor_slices(t)
ds = ds.batch(2)
ds = ds.rebatch(3)
for batch in ds:
print(batch)
Результат выполнения кода:
tf.Tensor(
[[1 2 3]
[4 5 6]
[7 8 9]], shape=(3, 3), dtype=int32)
tf.Tensor(
[[10 11 12]], shape=(1, 3), dtype=int32)
Смотрите также
-
метод
batch,
который группирует элементы датасета в батчи -
метод
unbatch,
который разбивает батчи на отдельные элементы -
метод
padded_batch,
который группирует элементы с дополнением до одинаковой формы -
класс
Dataset,
который представляет собой набор данных