tf.keras.utils.PyDataset
Базовый класс для определения параллельного набора данных с помощью кода Python.
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