Метод unbatch класса Dataset
Метод unbatch класса Dataset выполняет операцию, обратную методу batch.
Он разделяет элементы датасета, которые были сгруппированы в батчи,
обратно на отдельные элементы. Метод не принимает параметров и возвращает
новый объект Dataset, в котором каждый элемент соответствует
отдельному образцу из исходных батчей. Это полезно, когда нужно вернуться
к поэлементной обработке после пакетной.
Синтаксис
dataset.unbatch()
Пример
Давайте создадим датасет из тензора, сгруппируем элементы в батчи,
а затем применим метод unbatch, чтобы разделить их обратно:
import tensorflow as tf
dataset = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5, 6])
batched_dataset = dataset.batch(2)
for batch in batched_dataset:
print(batch.numpy())
unbatched_dataset = batched_dataset.unbatch()
for element in unbatched_dataset:
print(element.numpy())
Результат выполнения кода:
[1 2]
[3 4]
[5 6]
1
2
3
4
5
6
Пример
Давайте рассмотрим работу метода unbatch на двумерных тензорах,
где каждый элемент батча имеет несколько признаков:
import tensorflow as tf
dataset = tf.data.Dataset.from_tensor_slices([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
batched_dataset = dataset.batch(2)
for batch in batched_dataset:
print(batch.numpy())
unbatched_dataset = batched_dataset.unbatch()
for element in unbatched_dataset:
print(element.numpy())
Результат выполнения кода:
[[1 2 3]
[4 5 6]]
[[7 8 9]]
[1 2 3]
[4 5 6]
[7 8 9]
Пример
Давайте проверим, что метод unbatch восстанавливает исходное
количество элементов в датасете:
import tensorflow as tf
dataset = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
batched_dataset = dataset.batch(2)
unbatched_dataset = batched_dataset.unbatch()
print("Original cardinality:", dataset.cardinality().numpy())
print("Batched cardinality:", batched_dataset.cardinality().numpy())
print("Unbatched cardinality:", unbatched_dataset.cardinality().numpy())
Результат выполнения кода:
Original cardinality: 5
Batched cardinality: 3
Unbatched cardinality: 5
Смотрите также
-
метод
batch,
который объединяет элементы датасета в батчи -
класс
Dataset,
который представляет коллекцию элементов -
метод
from_tensor_slices,
который создает датасет из тензора -
метод
rebatch,
который перегруппировывает элементы в батчи