Spec-Zone.ru › TensorFlow

tf.keras.utils.PyDataset

Базовый класс для определения параллельного набора данных с помощью кода Python.

Просмотр псевдонимов

Основные псевдонимы

tf.keras.utils.Sequence

tf.keras.utils.PyDataset(
    workers=1, use_multiprocessing=False, max_queue_size=10
)

Каждый PyDataset должен реализовывать методы __getitem__() и __len__(). Если вы хотите изменить свой набор данных между эпохами, вы можете дополнительно реализовать on_epoch_end(). Метод __getitem__() должен возвращать полный пакет (а не один образец), а метод __len__ должен возвращать количество пакетов в наборе данных (а не количество образцов).

Аргументы
workers Количество потоков для использования в многопоточном или многопроцессном режиме.
use_multiprocessing Использовать ли многопроцессорность Python для параллелизма. Установка этого значения в True означает, что ваш набор данных будет дублирован в нескольких разветвлённых процессах. Это необходимо для получения вычислительных (а не ввода-вывода) преимуществ от параллелизма. Однако это значение можно установить только в True, если ваш набор данных можно безопасно сохранить.
max_queue_size Максимальное количество пакетов для хранения в очереди при итерации по набору данных в многопоточном или многопроцессном режиме. Уменьшите это значение для уменьшения использования процессорной памяти вашим набором данных. По умолчанию равно 10.

Примечания:

  • PyDataset — более безопасный способ использования многопроцессорности. Эта структура гарантирует, что модель будет обучаться только один раз на каждом образце за эпоху, что не так в случае с Python-генераторами.
  • Аргументы workers, use_multiprocessing и max_queue_size предназначены для настройки того, как fit() использует параллелизм для итерации по набору данных. Они не используются классом PyDataset напрямую. При ручном переборе PyDataset параллелизм не применяется.

Пример:

from skimage.io import imread
from skimage.transform import resize
import numpy as np
import math

# Here, `x_set` is list of path to the images
# and `y_set` are the associated classes.

class CIFAR10PyDataset(keras.utils.PyDataset):

    def __init__(self, x_set, y_set, batch_size, **kwargs):
        super().__init__(**kwargs)
        self.x, self.y = x_set, y_set
        self.batch_size = batch_size

    def __len__(self):
        # Return number of batches.
        return math.ceil(len(self.x) / self.batch_size)

    def __getitem__(self, idx):
        # Return x, y for batch idx.
        low = idx * self.batch_size
        # Cap upper bound at array length; the last batch may be smaller
        # if the total number of items is not a multiple of batch size.
        high = min(low + self.batch_size, len(self.x))
        batch_x = self.x[low:high]
        batch_y = self.y[low:high]

        return np.array([
            resize(imread(file_name), (200, 200))
               for file_name in batch_x]), np.array(batch_y)
Атрибуты
max_queue_size
num_batches Количество пакетов в PyDataset.
use_multiprocessing
workers

Методы

on_epoch_end

Просмотреть исходный код

on_epoch_end()

Метод, вызываемый в конце каждой эпохи.

__getitem__

Просмотреть исходный код

__getitem__(
    index
)

Получает пакет в позиции index.

Аргументы
index Позиция пакета в PyDataset.
Возвращаемые значения
Пакет

© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/api_docs/python/tf/keras/utils/PyDataset

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API