Класс Dataset
Класс Dataset - это базовый класс для работы с наборами данных
в TensorFlow. Он представляет собой последовательность элементов,
по которой можно итерироваться, преобразовывать её с помощью методов
(например, map, filter, batch) и эффективно
подавать данные в модель. Экземпляр класса обычно не создают
напрямую через конструктор, а получают с помощью фабричных методов
from_tensor_slices, from_tensors, from_generator,
range и других. Каждый элемент датасета - это тензор или
кортеж тензоров, описываемых атрибутом element_spec.
Синтаксис
tf.data.Dataset.from_tensor_slices(tensors)
tf.data.Dataset.from_tensors(tensors)
tf.data.Dataset.range(*args)
Пример
Давайте создадим датасет из списка чисел и выведем его элементы:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
for elem in ds:
print(elem)
Результат выполнения кода:
tf.Tensor(1, shape=(), dtype=int32)
tf.Tensor(2, shape=(), dtype=int32)
tf.Tensor(3, shape=(), dtype=int32)
tf.Tensor(4, shape=(), dtype=int32)
tf.Tensor(5, shape=(), dtype=int32)
Пример
Давайте применим к датасету преобразование map, чтобы
умножить каждый элемент на два:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
ds = ds.map(lambda x: x * 2)
for elem in ds:
print(elem)
Результат выполнения кода:
tf.Tensor(2, shape=(), dtype=int32)
tf.Tensor(4, shape=(), dtype=int32)
tf.Tensor(6, shape=(), dtype=int32)
tf.Tensor(8, shape=(), dtype=int32)
tf.Tensor(10, shape=(), dtype=int32)
Пример
Давайте сгруппируем элементы датасета в батчи по два
с помощью метода batch:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
ds = ds.batch(2)
for elem in ds:
print(elem)
Результат выполнения кода:
tf.Tensor([1 2], shape=(2,), dtype=int32)
tf.Tensor([3 4], shape=(2,), dtype=int32)
tf.Tensor([5], shape=(1,), dtype=int32)
Пример
Давайте создадим датасет из двумерного тензора и посмотрим
на его спецификацию через атрибут element_spec:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices(
tf.constant([[1, 2, 3], [4, 5, 6]])
)
print(ds.element_spec)
Результат выполнения кода:
TensorSpec(shape=(3,), dtype=tf.int32, name=None)
Смотрите также
-
метод
from_tensor_slices,
который создает датасет из тензоров -
метод
from_tensors,
который создает датасет из одного тензора -
метод
batch,
который группирует элементы датасета в батчи -
метод
map,
который применяет функцию к каждому элементу