РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
478 of 824 menu

Метод batch класса Dataset

Метод batch класса Dataset объединяет последовательные элементы датасета в пакеты. Первым параметром метод принимает размер пакета (целое число или tf.int64). Вторым необязательным параметром можно передать drop_remainder - логическое значение, которое указывает, отбрасывать ли последний неполный пакет. Метод возвращает новый датасет, элементами которого являются пакеты.

Синтаксис

Dataset.batch(batch_size, drop_remainder=False)

Пример

Давайте создадим датасет из тензора и сгруппируем его элементы в пакеты по 2:

import tensorflow as tf ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5]) ds = ds.batch(2) for batch in ds: print(batch)

Результат выполнения кода:

tf.Tensor([1 2], shape=(2,), dtype=int32) tf.Tensor([3 4], shape=(2,), dtype=int32) tf.Tensor([5], shape=(1,), dtype=int32)

Пример

Давайте отбросим последний неполный пакет с помощью параметра drop_remainder:

import tensorflow as tf ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5]) ds = ds.batch(2, drop_remainder=True) for batch in ds: print(batch)

Результат выполнения кода:

tf.Tensor([1 2], shape=(2,), dtype=int32) tf.Tensor([3 4], shape=(2,), dtype=int32)

Пример

Давайте сгруппируем в пакеты двумерные элементы датасета:

import tensorflow as tf t = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) ds = tf.data.Dataset.from_tensor_slices(t) ds = ds.batch(2) for batch in ds: print(batch)

Результат выполнения кода:

tf.Tensor( [[1 2 3] [4 5 6]], shape=(2, 3), dtype=int32) tf.Tensor([[7 8 9]], shape=(1, 3), dtype=int32)

Пример

Давайте применим batch вместе с shuffle и repeat для обучения модели:

import tensorflow as tf tf.random.set_seed(0) ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5]) ds = ds.shuffle(5).batch(2).repeat(2) for batch in ds: print(batch)

Результат выполнения кода:

tf.Tensor([3 1], shape=(2,), dtype=int32) tf.Tensor([5 2], shape=(2,), dtype=int32) tf.Tensor([4], shape=(1,), dtype=int32) tf.Tensor([1 3], shape=(2,), dtype=int32) tf.Tensor([2 5], shape=(2,), dtype=int32) tf.Tensor([4], shape=(1,), dtype=int32)

Смотрите также

  • класс Dataset,
    который представляет набор данных
  • метод padded_batch,
    который группирует элементы с выравниванием длины
  • метод unbatch,
    который разбивает пакеты на отдельные элементы
  • метод rebatch,
    который изменяет размер пакетов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить