torch.utils.data
В основе PyTorch-утилиты загрузки данных лежит класс torch.utils.data.DataLoader. Он представляет собой итерируемый объект Python над набором данных, поддерживающий
- стили наборов данных map и iterable,
- настраиваемый порядок загрузки данных,
- автоматическое формирование пакетов,
- загрузку данных с использованием одного или нескольких процессов,
- автоматическое закрепление памяти.
Эти параметры настраиваются аргументами конструктора класса DataLoader, у которого есть сигнатура:
DataLoader(dataset, batch_size=1, shuffle=False, sampler=None,
batch_sampler=None, num_workers=0, collate_fn=None,
pin_memory=False, drop_last=False, timeout=0,
worker_init_fn=None, *, prefetch_factor=2,
persistent_workers=False)
В следующих разделах подробно описаны эффекты и использование этих параметров.
Типы наборов данных
Самый важный аргумент конструктора DataLoader — это dataset, указывающий на объект набора данных для загрузки данных. PyTorch поддерживает два разных типа наборов данных:
Наборы данных типа map
Набор данных типа map реализует протоколы __getitem__() и __len__(), и представляет собой отображение (возможно, не целых) индексов/ключей на образцы данных.
Например, такой набор данных, при обращении с помощью dataset[idx], может считать idx-й образ и соответствующую метку из папки на диске.
См. Dataset для получения более подробной информации.
Наборы данных типа iterable
Набор данных типа iterable — это экземпляр подкласса IterableDataset, реализующий протокол __iter__(), и представляет собой итерируемый объект над образцами данных. Этот тип наборов данных особенно подходит для случаев, когда случайный доступ к данным является дорогостоящим или даже маловероятным, и размер пакета зависит от извлечённых данных.
Например, такой набор данных, при вызове iter(dataset), может вернуть поток данных, читаемых из базы данных, удаленного сервера или даже логов, генерируемых в режиме реального времени.
См. IterableDataset для получения более подробной информации.
Примечание
При использовании IterableDataset с многопроцессорной загрузкой данных. Тот же объект набора данных дублируется в каждом процессе-рабочем, и поэтому копии должны быть сконфигурированы по-разному, чтобы избежать дублирования данных. См. документацию IterableDataset для получения информации о том, как этого добиться.
Порядок загрузки данных и выборщик
Для наборов данных типа iterable порядок загрузки данных полностью определяется пользователем-определённой итерацией. Это позволяет более просто реализовать чтение блоками и динамический размер пакета (например, с помощью возврата запакованного образца в каждый момент).
Остальная часть этого раздела относится к наборам данных типа map. Используются классы torch.utils.data.Sampler для определения последовательности индексов/ключей, используемых при загрузке данных. Они представляют собой итерируемые объекты над индексами наборов данных. Например, в общем случае с градиентным спуском со случайной выборкой (SGD), Sampler может случайным образом переупорядочить список индексов и выводить каждый по одному, или выводить небольшое количество из них для мини-пакетного SGD.
Последовательный или перемешанный выборщик будет автоматически сконструирован на основе аргумента shuffle для DataLoader. В качестве альтернативы, пользователи могут использовать аргумент sampler для указания пользовательского объекта выборщика Sampler, который в каждый момент возвращает следующий индекс/ключ для извлечения.
Пользовательский Sampler, возвращающий список индексов пакета за раз, может быть передан в качестве аргумента batch_sampler. Автоматическое формирование пакетов также может быть включено с помощью аргументов batch_size и drop_last. Дополнительные сведения об этом см. в следующем разделе .
Примечание
Ни sampler ни batch_sampler несовместимы с наборами данных типа iterable, поскольку такие наборы данных не имеют понятия о ключе или индексе.
Загрузка данных пакетами и без пакетов
DataLoader поддерживает автоматическое объединение отдельных извлеченных образцов данных в пакеты с помощью аргументов batch_size, drop_last, batch_sampler, и collate_fn (с функцией по умолчанию).
Автоматическое формирование пакетов (по умолчанию)
Это наиболее распространенный случай, и соответствует извлечению мини-пакета данных и объединению их в образцы пакетов, т.е. содержащие тензоры с одним измерением, являющимся размерностью пакета (обычно первым).
Если batch_size (значение по умолчанию 1) не равно None, то загрузчик данных возвращает образцы пакетов вместо отдельных образцов. Аргументы batch_size и drop_last используются для указания того, как загрузчик данных получает пакеты ключей набора данных. Для наборов данных типа map пользователи могут альтернативно указать batch_sampler, которое возвращает список ключей за один раз.
Примечание
Аргументы batch_size и drop_last по существу используются для создания batch_sampler из sampler. Для наборов данных типа map sampler либо предоставляется пользователем, либо создаётся на основе аргумента shuffle. Для наборов данных типа iterable sampler является фиктивным бесконечным.
Примечание
При извлечении из наборов данных типа iterable с многопроцессорной обработкой, аргумент drop_last отбрасывает последний неполный пакет в каждом наборе данных копии каждого рабочего процесса.
После извлечения списка образцов с помощью индексов из выборщика, используется функция, переданная в качестве аргумента collate_fn, для объединения списков образцов в пакеты.
В этом случае загрузка из набора данных типа map примерно эквивалентна:
for indices in batch_sampler:
yield collate_fn([dataset[i] for i in indices])
а загрузка из набора данных типа iterable примерно эквивалентна:
dataset_iter = iter(dataset)
for indices in batch_sampler:
yield collate_fn([next(dataset_iter) for _ in indices])
Пользовательская функция collate_fn может использоваться для настройки объединения, например, для заполнения последовательных данных максимальной длиной пакета. См. этот раздел для получения дополнительной информации о collate_fn.
Отключение автоматического формирования пакетов
В некоторых случаях пользователи могут захотеть вручную обрабатывать формирование пакетов в коде набора данных или просто загружать отдельные образцы. Например, это может быть более эффективным для непосредственной загрузки данных пакетами (например, массового чтения из базы данных или чтения непрерывных блоков памяти), размер пакета зависит от данных или программа разработана для работы с отдельными образцами. В таких сценариях, вероятно, лучше не использовать автоматическое формирование пакетов (где collate_fn используется для объединения образцов), а позволить загрузчику данных напрямую возвращать каждый член объекта dataset.
Когда и batch_size и batch_sampler равны None (значение по умолчанию для batch_sampler уже None), автоматическое формирование пакетов отключено. Каждый образец, полученный от dataset обрабатывается с помощью функции, переданной в качестве аргумента collate_fn.
При отключении автоматического формирования пакетов функция collate_fn по умолчанию просто преобразует массивы NumPy в тензоры PyTorch и оставляет всё остальное без изменений.
В этом случае загрузка из набора данных типа map примерно эквивалентна:
for index in sampler:
yield collate_fn(dataset[index])
а загрузка из набора данных типа iterable примерно эквивалентна:
for data in iter(dataset):
yield collate_fn(data)
См. этот раздел для получения дополнительной информации о collate_fn.
Работа с collate_fn
Использование collate_fn немного отличается, когда автоматическое формирование пакетов включено или выключено.
При отключенном автоматическом формировании пакетов, collate_fn вызывается с каждым отдельным образцом данных, и результат возвращается итератором загрузчика данных. В этом случае, по умолчанию collate_fn просто преобразует массивы NumPy в тензоры PyTorch.
При включенном автоматическом формировании пакетов, collate_fn вызывается со списком образцов данных в каждый момент. Ожидается, что он объединит входные образцы в пакет для возврата итератором загрузчика данных. Остальная часть этого раздела описывает поведение функции collate_fn по умолчанию (default_collate()).
Например, если каждый образец данных состоит из 3-канального изображения и целочисленной метки класса, т. е., каждый элемент набора данных возвращает кортеж (image, class_index), по умолчанию collate_fn собирает список таких кортежей в один кортеж с тензором изображения пакета и тензором метки класса пакета. В частности, по умолчанию collate_fn обладает следующими свойствами:
- Он всегда добавляет новую размерность в качестве размерности пакета.
- Он автоматически преобразует массивы NumPy и числовые значения Python в тензоры PyTorch.
- Он сохраняет структуру данных, например, если каждый образец — это словарь, он выводит словарь с тем же набором ключей, но с тензорами пакета в качестве значений (или списками, если значения нельзя преобразовать в тензоры). То же самое для
list,tuple,namedtupleи т. д.
Пользователи могут использовать настраиваемые collate_fn для достижения пользовательского пакетного формирования, например, для пакетного формирования по размерности, отличной от первой, для дополнения последовательностей различной длины или для добавления поддержки пользовательских типов данных.
Если у вас возникнет ситуация, когда размеры или тип вывода DataLoader отличаются от ожидаемых, вы можете проверить свой collate_fn.
Загрузка данных с использованием одного и нескольких процессов
A DataLoader по умолчанию использует загрузку данных с использованием одного процесса.
В рамках процесса Python блокировка глобального интерпретатора (GIL) препятствует настоящему полному параллелизму кода Python по потокам. Для предотвращения блокировки вычислений кода с загрузкой данных PyTorch предоставляет простой способ переключения на загрузку данных с использованием нескольких процессов, просто установив аргумент num_workers в положительное целое число.
Загрузка данных с использованием одного процесса (по умолчанию)
В этом режиме извлечение данных выполняется в том же процессе, в котором инициализируется DataLoader. Поэтому загрузка данных может заблокировать вычисления. Однако этот режим может быть предпочтительным, когда ресурсы, используемые для обмена данными между процессами (например, общая память, дескрипторы файлов), ограничены или когда весь набор данных небольшой и может быть загружен полностью в память. Кроме того, загрузка с использованием одного процесса часто показывает более читаемые трассировки ошибок и, следовательно, полезна для отладки.
Загрузка данных с использованием нескольких процессов
Установка аргумента num_workers в положительное целое число включит загрузку данных с использованием нескольких процессов с указанным количеством рабочих процессов загрузчика.
Предупреждение
После нескольких итераций рабочие процессы загрузчика будут потреблять столько же памяти процессора, сколько и родительский процесс, для всех объектов Python в родительском процессе, к которым обращаются рабочие процессы. Это может быть проблематично, если набор данных содержит много данных (например, вы загружаете очень большой список имен файлов во время создания набора данных) и/или вы используете много рабочих процессов (общее использование памяти составляет number of workers * size of parent process). Самым простым способом решения этой проблемы является замена объектов Python на представления без подсчёта ссылок, такие как объекты Pandas, Numpy или PyArrow. Более подробную информацию о причинах возникновения этой проблемы и примеры кода для решения этих проблем см. в проблеме #13246.
В этом режиме каждый раз, когда создается итератор DataLoader (например, при вызове enumerate(dataloader)), создаются num_workers рабочих процессов. На этом этапе dataset, collate_fn, и worker_init_fn передаются каждому рабочему процессу, где они используются для инициализации и извлечения данных. Это означает, что доступ к набору данных вместе с его внутренней работой ввода-вывода, преобразованиями (включая collate_fn) выполняется в рабочем процессе.
torch.utils.data.get_worker_info() возвращает различные полезные сведения в рабочем процессе (включая идентификатор рабочего процесса, копию набора данных, начальное значение, и т. д.), и возвращает None в основном процессе. Пользователи могут использовать эту функцию в коде набора данных и/или worker_init_fn для индивидуальной настройки каждой копии набора данных и для определения того, выполняется ли код в рабочем процессе. Например, это может быть особенно полезно при фрагментации набора данных.
Для наборов данных типа «карта» основной процесс генерирует индексы с использованием sampler и отправляет их рабочим процессам. Таким образом, любая случайная перестановка выполняется в основном процессе, который управляет загрузкой, назначая индексы для загрузки.
Для наборов данных типа «итерируемый» каждый рабочий процесс получает копию объекта dataset, поэтому примитивная загрузка с использованием нескольких процессов часто приводит к дублированию данных. Используя torch.utils.data.get_worker_info() и/или worker_init_fn, пользователи могут настроить каждую копию независимо. (См. документацию IterableDataset для получения инструкций о том, как это сделать.) По тем же причинам в загрузке с использованием нескольких процессов аргумент drop_last отбрасывает последнюю неполную пачку каждой копии набора данных типа «итерируемый» рабочего процесса.
Рабочие процессы закрываются по достижении конца итерации или при сборке мусора итератора.
Предупреждение
В целом не рекомендуется возвращать тензоры CUDA при загрузке с использованием нескольких процессов из-за множества тонкостей в использовании CUDA и обмене тензорами CUDA в параллельных вычислениях (см. CUDA в параллельных вычислениях). Вместо этого рекомендуется использовать автоматическое закрепление памяти (т. е. установку pin_memory=True), что обеспечивает быстрый обмен данными с графическими процессорами, поддерживающими CUDA.
Платформозависимое поведение
Поскольку рабочие процессы полагаются на Python multiprocessing, поведение запуска рабочих процессов отличается на Windows от Unix.
- В Unix
fork()— это метод запуска по умолчаниюmultiprocessing. Используяfork(), дочерние рабочие процессы, как правило, могут напрямую получать доступ к функциямdatasetи Python-аргументам через клонированное адресное пространство. - В Windows или MacOS
spawn()— это метод запуска по умолчаниюmultiprocessing. Используяspawn(), запускается другой интерпретатор, который выполняет ваш основной скрипт, за которым следует внутренняя функция рабочего процесса, которая получаетdataset,collate_fnи другие аргументы через сериализациюpickle.
Эта отдельная сериализация означает, что для обеспечения совместимости с Windows при использовании загрузки данных с использованием нескольких процессов вам необходимо выполнить два шага:
- Упакуйте большую часть кода вашего основного скрипта в блок
if __name__ == '__main__':, чтобы убедиться, что он не выполняется повторно (что, скорее всего, вызовет ошибку) при запуске каждого рабочего процесса. Вы можете поместить сюда логику создания набора данных и экземпляраDataLoader, так как ее не нужно выполнять повторно в рабочих процессах. - Убедитесь, что любой пользовательский код
collate_fn,worker_init_fnилиdatasetобъявлен как определения верхнего уровня, за пределами проверки__main__. Это гарантирует, что они доступны в рабочих процессах. (Это необходимо, поскольку функции сериализуются только как ссылки, а не какbytecode.)
Случайность при загрузке данных с использованием нескольких процессов
По умолчанию для каждого рабочего процесса устанавливается значение генератора случайных чисел PyTorch, равное base_seed + worker_id, где base_seed — это целое число, сгенерированное основным процессом с использованием его генератора случайных чисел (тем самым, обязательно потребляя состояние генератора случайных чисел) или указанное generator. Однако значения семян для других библиотек могут дублироваться при инициализации рабочих процессов, в результате чего каждый рабочий процесс будет возвращать одинаковые случайные числа. (См. этот раздел в разделе вопросов и ответов.)
В worker_init_fn, вы можете получить доступ к установленному значению семени PyTorch для каждого рабочего процесса с помощью либо torch.utils.data.get_worker_info().seed, либо torch.initial_seed() и использовать его для задания семян других библиотек перед загрузкой данных.
Закрепление памяти
Копирование с хоста на графический процессор значительно быстрее, когда оно происходит из закреплённой (заблокированной в странице) памяти. Более подробную информацию о том, когда и как использовать закреплённую память, см. в Использование буферов закрепленной памяти.
При загрузке данных передача pin_memory=True в DataLoader автоматически помещает извлеченные тензоры данных в закреплённую память, тем самым обеспечивая более быстрый обмен данными с графическими процессорами, поддерживающими CUDA.
Логика закрепления памяти по умолчанию распознает только тензоры и отображения и итерируемые объекты, содержащие тензоры. По умолчанию, если логика закрепления встречает пакет пользовательского типа (что произойдёт, если у вас есть collate_fn функция, возвращающая пакет пользовательского типа) или если каждый элемент вашего пакета имеет пользовательский тип, логика закрепления не распознает их и вернёт этот пакет (или эти элементы) без закрепления памяти. Для включения закрепления памяти для пользовательского пакета или типа данных определите метод pin_memory() для вашего(их) пользовательского(их) типа(ов).
См. пример ниже.
Пример:
class SimpleCustomBatch:
def __init__(self, data):
transposed_data = list(zip(*data))
self.inp = torch.stack(transposed_data[0], 0)
self.tgt = torch.stack(transposed_data[1], 0)
# custom memory pinning method on custom type
def pin_memory(self):
self.inp = self.inp.pin_memory()
self.tgt = self.tgt.pin_memory()
return self
def collate_wrapper(batch):
return SimpleCustomBatch(batch)
inps = torch.arange(10 * 5, dtype=torch.float32).view(10, 5)
tgts = torch.arange(10 * 5, dtype=torch.float32).view(10, 5)
dataset = TensorDataset(inps, tgts)
loader = DataLoader(dataset, batch_size=2, collate_fn=collate_wrapper,
pin_memory=True)
for batch_ndx, sample in enumerate(loader):
print(sample.inp.is_pinned())
print(sample.tgt.is_pinned())
-
class torch.utils.data.DataLoader(dataset, batch_size=1, shuffle=None, sampler=None, batch_sampler=None, num_workers=0, collate_fn=None, pin_memory=False, drop_last=False, timeout=0, worker_init_fn=None, multiprocessing_context=None, generator=None, *, prefetch_factor=None, persistent_workers=False, pin_memory_device='')[source] -
Загрузчик данных. Объединяет набор данных и выборщик, обеспечивая возможность итерирования по заданному набору данных.
Загрузчик данных
DataLoaderподдерживает наборы данных с стилем «отображение» и «итерируемый» с загрузкой с использованием одного или нескольких процессов, настраиваемым порядком загрузки и необязательным автоматическим формированием батчей (коллекцией) и закреплением памяти.См. страницу документации
torch.utils.dataдля получения более подробной информации.- Параметры
-
- dataset (Dataset) – набор данных, из которого загружаются данные.
-
batch_size (int, необязательно) – количество образцов в батче (по умолчанию:
1). -
shuffle (bool, необязательно) – устанавливается в
Trueдля перетасовки данных на каждом эпохе (по умолчанию:False). -
sampler (Sampler или Итерируемый, необязательно) – определяет стратегию выбора образцов из набора данных. Может быть любым
Iterableс реализованным__len__. Если указан,shuffleне должен быть указан. -
batch_sampler (Sampler или Итерируемый, необязательно) – аналогично
sampler, но возвращает батч индексов за раз. Взаимоисключающее сbatch_size,shuffle,sampler, иdrop_last. -
num_workers (int, необязательно) – количество подпроцессов для загрузки данных.
0означает, что данные будут загружаться в основном процессе. (по умолчанию:0). - collate_fn (Вызываемый объект, необязательно) – объединяет список образцов в мини-батчи тензоров. Используется при использовании пакетной загрузки из набора данных со стилем «отображение».
-
pin_memory (bool, необязательно) – Если
True, загрузчик данных скопирует тензоры в закреплённую память устройства/CUDA перед их возвращением. Если ваши элементы данных являются пользовательским типом, или вашcollate_fnвозвращает батч пользовательского типа, см. пример ниже. -
drop_last (bool, необязательно) – устанавливается в
Trueдля отбрасывания последнего неполного батча, если размер набора данных не делится на размер батча. ЕслиFalseи размер набора данных не делится на размер батча, то последний батч будет меньше. (по умолчанию:False). -
timeout (числовое, необязательно) – если положительно, время ожидания для сбора батча от рабочих процессов. Должно быть всегда неотрицательным. (по умолчанию:
0). -
worker_init_fn (Вызываемый объект, необязательно) – Если не
None, это будет вызвано в каждом подпроцессе рабочего с идентификатором рабочего процесса (целое число в[0, num_workers - 1]) в качестве входных данных после инициализации и перед загрузкой данных. (по умолчанию:None). -
multiprocessing_context (str или multiprocessing.context.BaseContext, необязательно) – Если
None, будет использоваться по умолчанию контекст multiprocessing вашей операционной системы. (по умолчанию:None). -
generator (torch.Generator, необязательно) – Если не
None, этот генератор ПСП будет использоваться RandomSampler для генерации случайных индексов и многопроцессностью для генерацииbase_seedдля рабочих процессов. (по умолчанию:None). -
prefetch_factor (int, необязательно, ключевой аргумент) – Количество батчей, загружаемых заранее каждым рабочим процессом.
2означает, что всего будет 2 * num_workers батча предварительной загрузки по всем рабочим процессам. (Значение по умолчанию зависит от установленного значения для num_workers. Если значение num_workers=0, значение по умолчаниюNone. В противном случае, если значениеnum_workers > 0значение по умолчанию2). -
persistent_workers (bool, необязательно) – Если
True, загрузчик данных не будет завершать процессы рабочих процессов после однократного потребления набора данных. Это позволяет сохранить рабочие процессыDatasetэкземпляры живыми. (по умолчанию:False). -
pin_memory_device (str, необязательно) – устройство для
pin_memoryв случае, еслиpin_memoryравенTrue.
Предупреждение
Если используется метод запуска
spawn,worker_init_fnне может быть несериализуемым объектом, например, лямбда-функцией. См. Рекомендации по работе с многопроцессорностью для получения дополнительной информации о работе с многопроцессорностью в PyTorch.Предупреждение
Эвристика
len(dataloader)основана на длине используемого выборщика. КогдаdatasetявляетсяIterableDataset, он вместо этого возвращает оценку, основанную наlen(dataset) / batch_size, с соответствующим округлением в зависимости отdrop_last, независимо от конфигураций многопроцессорной загрузки. Это наилучшая оценка, которую может сделать PyTorch, так как PyTorch доверяет коду пользователяdatasetдля правильной обработки многопроцессорной загрузки для предотвращения дублирования данных.Однако, если фрагментация приводит к тому, что несколько рабочих процессов имеют неполные последние батчи, эта оценка все равно может быть неточной, потому что (1) в противном случае полный батч может быть разделен на несколько, и (2) при установке
drop_lastможет быть отброшено более одного батча образцов. К сожалению, PyTorch не может обнаружить такие случаи в общем случае.См. Типы наборов данных для получения более подробной информации об этих двух типах наборов данных и о том, как
IterableDatasetвзаимодействует с Многопроцессорная загрузка данных.Предупреждение
См. Воспроизводимость, Рабочие процессы загрузчика данных возвращают идентичные случайные числа и Случайность в многопроцессорной загрузке данных для вопросов, связанных со случайным началом.
-
class torch.utils.data.Dataset(*args, **kwds)[source] -
Абстрактный класс, представляющий набор данных
Dataset.Все наборы данных, представляющие отображение ключей на образцы данных, должны наследоваться от него. Все подклассы должны переопределять
__getitem__(), поддерживая получение образца данных для заданного ключа. Подклассы также могут необязательно переопределять__len__(), которое ожидается, что вернёт размер набора данных многими реализациямиSamplerи опциями по умолчаниюDataLoader. Подклассы также могут необязательно реализовывать__getitems__(), для ускорения загрузки батчей образцов. Этот метод принимает список индексов образцов батча и возвращает список образцов.Примечание
DataLoaderпо умолчанию строит выборщик индексов, который возвращает целочисленные индексы. Чтобы заставить его работать с набором данных со стилем «отображение» с индексами/ключами, не являющимися целыми числами, необходимо предоставить пользовательский выборщик.
-
class torch.utils.data.IterableDataset(*args, **kwds)[source] -
Итерируемый набор данных.
Все наборы данных, представляющие итерируемый набор образцов данных, должны быть его подклассом. Такая форма наборов данных особенно полезна, когда данные поступают из потока.
Все подклассы должны перезаписать
__iter__(), которая вернёт итератор образцов в этом наборе данных.При использовании подкласса с
DataLoader, каждый элемент в наборе данных будет выводиться из итератораDataLoader. Когдаnum_workers > 0, каждый процесс-воркер будет иметь различную копию объекта набора данных, поэтому часто желательно настроить каждую копию независимо, чтобы избежать дублирования данных, возвращаемых воркерами.get_worker_info(), когда вызывается в процессе воркера, возвращает информацию о воркере. Она может использоваться в методе набора данных__iter__()или в параметреDataLoader‘sworker_init_fnдля изменения поведения каждой копии.Пример 1: распределение рабочей нагрузки между всеми воркерами в
__iter__():>>> class MyIterableDataset(torch.utils.data.IterableDataset): ... def __init__(self, start, end): ... super(MyIterableDataset).__init__() ... assert end > start, "this example code only works with end >= start" ... self.start = start ... self.end = end ... ... def __iter__(self): ... worker_info = torch.utils.data.get_worker_info() ... if worker_info is None: # single-process data loading, return the full iterator ... iter_start = self.start ... iter_end = self.end ... else: # in a worker process ... # split workload ... per_worker = int(math.ceil((self.end - self.start) / float(worker_info.num_workers))) ... worker_id = worker_info.id ... iter_start = self.start + worker_id * per_worker ... iter_end = min(iter_start + per_worker, self.end) ... return iter(range(iter_start, iter_end)) ... >>> # should give same set of data as range(3, 7), i.e., [3, 4, 5, 6]. >>> ds = MyIterableDataset(start=3, end=7) >>> # Single-process loading >>> print(list(torch.utils.data.DataLoader(ds, num_workers=0))) [tensor([3]), tensor([4]), tensor([5]), tensor([6])] >>> # Mult-process loading with two worker processes >>> # Worker 0 fetched [3, 4]. Worker 1 fetched [5, 6]. >>> print(list(torch.utils.data.DataLoader(ds, num_workers=2))) [tensor([3]), tensor([5]), tensor([4]), tensor([6])] >>> # With even more workers >>> print(list(torch.utils.data.DataLoader(ds, num_workers=12))) [tensor([3]), tensor([5]), tensor([4]), tensor([6])]
Пример 2: распределение рабочей нагрузки между всеми воркерами с использованием
worker_init_fn:>>> class MyIterableDataset(torch.utils.data.IterableDataset): ... def __init__(self, start, end): ... super(MyIterableDataset).__init__() ... assert end > start, "this example code only works with end >= start" ... self.start = start ... self.end = end ... ... def __iter__(self): ... return iter(range(self.start, self.end)) ... >>> # should give same set of data as range(3, 7), i.e., [3, 4, 5, 6]. >>> ds = MyIterableDataset(start=3, end=7) >>> # Single-process loading >>> print(list(torch.utils.data.DataLoader(ds, num_workers=0))) [3, 4, 5, 6] >>> >>> # Directly doing multi-process loading yields duplicate data >>> print(list(torch.utils.data.DataLoader(ds, num_workers=2))) [3, 3, 4, 4, 5, 5, 6, 6] >>> # Define a `worker_init_fn` that configures each dataset copy differently >>> def worker_init_fn(worker_id): ... worker_info = torch.utils.data.get_worker_info() ... dataset = worker_info.dataset # the dataset copy in this worker process ... overall_start = dataset.start ... overall_end = dataset.end ... # configure the dataset to only process the split workload ... per_worker = int(math.ceil((overall_end - overall_start) / float(worker_info.num_workers))) ... worker_id = worker_info.id ... dataset.start = overall_start + worker_id * per_worker ... dataset.end = min(dataset.start + per_worker, overall_end) ... >>> # Mult-process loading with the custom `worker_init_fn` >>> # Worker 0 fetched [3, 4]. Worker 1 fetched [5, 6]. >>> print(list(torch.utils.data.DataLoader(ds, num_workers=2, worker_init_fn=worker_init_fn))) [3, 5, 4, 6] >>> # With even more workers >>> print(list(torch.utils.data.DataLoader(ds, num_workers=12, worker_init_fn=worker_init_fn))) [3, 4, 5, 6]
-
class torch.utils.data.TensorDataset(*tensors)[source] -
Набор данных, обертывающий тензоры.
Каждый образец будет извлекаться путём индексирования тензоров по первому измерению.
- Параметры
-
*tensors (Tensor) – тензоры, имеющие одинаковый размер первого измерения.
-
class torch.utils.data.StackDataset(*args, **kwargs)[source] -
Набор данных как объединение нескольких наборов данных.
Этот класс полезен для объединения различных частей сложных входных данных, представленных наборами данных.
Пример
>>> images = ImageDataset() >>> texts = TextDataset() >>> tuple_stack = StackDataset(images, texts) >>> tuple_stack[0] == (images[0], texts[0]) >>> dict_stack = StackDataset(image=images, text=texts) >>> dict_stack[0] == {'image': images[0], 'text': texts[0]}
-
class torch.utils.data.ConcatDataset(datasets)[source] -
Набор данных как конкатенация нескольких наборов данных.
Этот класс полезен для объединения различных существующих наборов данных.
- Параметры
-
datasets (последовательность) – Список наборов данных, подлежащих конкатенации
-
class torch.utils.data.ChainDataset(datasets)[source] -
Набор данных для цепочки нескольких
IterableDataset.Этот класс полезен для объединения различных существующих потоков данных наборов данных. Операция цепочки выполняется в режиме реального времени, поэтому конкатенация масштабных наборов данных с помощью этого класса будет эффективной.
- Параметры
-
datasets (итерируемый из IterableDataset) – наборы данных, которые должны быть объединены в цепочку
-
class torch.utils.data.Subset(dataset, indices)[source] -
Подмножество набора данных по указанным индексам.
- Параметры
-
- dataset (Dataset) – Весь набор данных
- indices (последовательность) – Индексы в полном наборе, выбранные для подмножества
-
torch.utils.data._utils.collate.collate(batch, *, collate_fn_map=None)[source] -
Общая функция объединения, которая обрабатывает тип элементов внутри каждого пакета и открывает реестр функций для обработки определённых типов элементов.
default_collate_fn_mapпредоставляет стандартные функции объединения для тензоров, массивов NumPy, чисел и строк.- Параметры
-
- batch – один пакет, который нужно объединить
- collate_fn_map (Optional[Dict[Union[Type, Tuple[Type, ...]], Callable]]) – Необязательный словарь, сопоставляющий тип элемента с соответствующей функцией объединения. Если тип элемента отсутствует в этом словаре, эта функция пройдёт по каждому ключу словаря в порядке вставки, чтобы вызвать соответствующую функцию объединения, если тип элемента является подклассом ключа.
Примеры
>>> # Extend this function to handle batch of tensors >>> def collate_tensor_fn(batch, *, collate_fn_map): ... return torch.stack(batch, 0) >>> def custom_collate(batch): ... collate_map = {torch.Tensor: collate_tensor_fn} ... return collate(batch, collate_fn_map=collate_map) >>> # Extend `default_collate` by in-place modifying `default_collate_fn_map` >>> default_collate_fn_map.update({torch.Tensor: collate_tensor_fn})Примечание
Каждая функция объединения требует позиционного аргумента для пакета и ключевого аргумента для словаря функций объединения как
collate_fn_map.
-
torch.utils.data.default_collate(batch)[source] -
Функция, принимающая пакет данных и помещающая элементы внутри пакета в тензор с дополнительным внешним измерением — размер пакета. Точный тип вывода может быть
torch.Tensor, кортежемSequenceизtorch.Tensor, коллекциейtorch.Tensorили оставлен неизменным, в зависимости от типа входных данных. Она используется в качестве стандартной функции объединения, когдаbatch_sizeилиbatch_samplerопределены вDataLoader.Вот общее сопоставление типов входных данных (в зависимости от типа элемента в пакете) с типом выходных данных:
-
torch.Tensor->torch.Tensor(с добавленным внешним измерением размера пакета) - Массивы NumPy ->
torch.Tensor -
float->torch.Tensor -
int->torch.Tensor -
str->str(без изменений) -
bytes->bytes(без изменений) -
Mapping[K, V_i]->Mapping[K, default_collate([V_1, V_2, …])] -
NamedTuple[V1_i, V2_i, …]->NamedTuple[default_collate([V1_1, V1_2, …]), default_collate([V2_1, V2_2, …]), …] -
Sequence[V1_i, V2_i, …]->Sequence[default_collate([V1_1, V1_2, …]), default_collate([V2_1, V2_2, …]), …]
- Параметры
-
batch – один пакет, который нужно объединить
Примеры
>>> # Example with a batch of `int`s: >>> default_collate([0, 1, 2, 3]) tensor([0, 1, 2, 3]) >>> # Example with a batch of `str`s: >>> default_collate(['a', 'b', 'c']) ['a', 'b', 'c'] >>> # Example with `Map` inside the batch: >>> default_collate([{'A': 0, 'B': 1}, {'A': 100, 'B': 100}]) {'A': tensor([ 0, 100]), 'B': tensor([ 1, 100])} >>> # Example with `NamedTuple` inside the batch: >>> Point = namedtuple('Point', ['x', 'y']) >>> default_collate([Point(0, 0), Point(1, 1)]) Point(x=tensor([0, 1]), y=tensor([0, 1])) >>> # Example with `Tuple` inside the batch: >>> default_collate([(0, 1), (2, 3)]) [tensor([0, 2]), tensor([1, 3])] >>> # Example with `List` inside the batch: >>> default_collate([[0, 1], [2, 3]]) [tensor([0, 2]), tensor([1, 3])] >>> # Two options to extend `default_collate` to handle specific type >>> # Option 1: Write custom collate function and invoke `default_collate` >>> def custom_collate(batch): ... elem = batch[0] ... if isinstance(elem, CustomType): # Some custom condition ... return ... ... else: # Fall back to `default_collate` ... return default_collate(batch) >>> # Option 2: In-place modify `default_collate_fn_map` >>> def collate_customtype_fn(batch, *, collate_fn_map=None): ... return ... >>> default_collate_fn_map.update(CustoType, collate_customtype_fn) >>> default_collate(batch) # Handle `CustomType` automatically -
-
torch.utils.data.default_convert(data)[source] -
Функция, преобразующая каждый элемент массива NumPy в
torch.Tensor. Если входной параметр являетсяSequence,Collection, илиMapping, она пытается преобразовать каждый элемент внутри вtorch.Tensor. Если входной параметр не является массивом NumPy, он остается неизменным. Эта функция используется по умолчанию для объединения элементов при отсутствии явных параметровbatch_samplerиbatch_sizeвDataLoader.Общая схема преобразования типов входных данных в выходные данные аналогична схеме
default_collate(). Для получения более подробной информации см. описание там.- Параметры
-
data – отдельный элемент данных для преобразования
Примеры
>>> # Example with `int` >>> default_convert(0) 0 >>> # Example with NumPy array >>> default_convert(np.array([0, 1])) tensor([0, 1]) >>> # Example with NamedTuple >>> Point = namedtuple('Point', ['x', 'y']) >>> default_convert(Point(0, 0)) Point(x=0, y=0) >>> default_convert(Point(np.array(0), np.array(0))) Point(x=tensor(0), y=tensor(0)) >>> # Example with List >>> default_convert([np.array([0, 1]), np.array([2, 3])]) [tensor([0, 1]), tensor([2, 3])]
-
torch.utils.data.get_worker_info()[source] -
Возвращает информацию о рабочем процессе текущего итератора
DataLoader.При вызове в рабочем процессе возвращает объект, гарантированно имеющий следующие атрибуты:
-
id: идентификатор текущего рабочего процесса. -
num_workers: общее количество рабочих процессов. -
seed: значение случайного числа, установленное для текущего рабочего процесса. Это значение определяется генератором случайных чисел основного процесса и идентификатором рабочего процесса. Для получения дополнительной информации см. документациюDataLoader. -
dataset: копия объекта набора данных в этом процессе. Обратите внимание, что это будет другой объект в другом процессе, отличном от основного процесса.
При вызове в основном процессе возвращает
None.Примечание
При использовании в
worker_init_fn, переданном вDataLoader, этот метод может быть полезен для настройки каждого рабочего процесса по-разному, например, используяworker_idдля настройки объектаdatasetтак, чтобы он считывал только определенную часть разбиеного набора данных, или используяseedдля инициализации других библиотек, используемых в коде набора данных.- Возвращаемый тип
-
Optional[WorkerInfo]
-
-
torch.utils.data.random_split(dataset, lengths, generator=<torch._C.Generator object>)[source] -
Случайное разбиение набора данных на непересекающиеся подмножества заданной длины.
Если предоставлен список дробей, сумма которых равна 1, длины будут вычислены автоматически как floor(frac * len(dataset)) для каждой предоставленной дроби.
После вычисления длин, если остаются остатки, 1 единица распределяется в циклическом порядке по длинам до тех пор, пока не останется остатков.
Для воспроизводимых результатов можно зафиксировать генератор, например:
Пример
>>> generator1 = torch.Generator().manual_seed(42) >>> generator2 = torch.Generator().manual_seed(42) >>> random_split(range(10), [3, 7], generator=generator1) >>> random_split(range(30), [0.3, 0.3, 0.4], generator=generator2)
-
class torch.utils.data.Sampler(data_source=None)[source] -
Базовый класс для всех Sampler.
Каждый подкласс Sampler должен предоставить метод
__iter__(), обеспечивающий способ итерации по индексам или спискам индексов (пакетов) элементов набора данных, и метод__len__(), возвращающий длину возвращаемых итераторов.- Параметры
-
data_source (Dataset) – Этот аргумент не используется и будет удален в версии 2.2.0. Вы по-прежнему можете иметь пользовательскую реализацию, использующую его.
Пример
>>> class AccedingSequenceLengthSampler(Sampler[int]): >>> def __init__(self, data: List[str]) -> None: >>> self.data = data >>> >>> def __len__(self) -> int: >>> return len(self.data) >>> >>> def __iter__(self) -> Iterator[int]: >>> sizes = torch.tensor([len(x) for x in self.data]) >>> yield from torch.argsort(sizes).tolist() >>> >>> class AccedingSequenceLengthBatchSampler(Sampler[List[int]]): >>> def __init__(self, data: List[str], batch_size: int) -> None: >>> self.data = data >>> self.batch_size = batch_size >>> >>> def __len__(self) -> int: >>> return (len(self.data) + self.batch_size - 1) // self.batch_size >>> >>> def __iter__(self) -> Iterator[List[int]]: >>> sizes = torch.tensor([len(x) for x in self.data]) >>> for batch in torch.chunk(torch.argsort(sizes), len(self)): >>> yield batch.tolist()
Примечание
Метод
__len__()не строго обязателен дляDataLoader, но ожидается при любых вычислениях, связанных с длинойDataLoader.
-
class torch.utils.data.SequentialSampler(data_source)[source] -
Обработка элементов последовательно, всегда в одном и том же порядке.
- Параметры
-
data_source (Dataset) – набор данных для выборки
-
class torch.utils.data.RandomSampler(data_source, replacement=False, num_samples=None, generator=None)[source] -
Выборка элементов случайным образом. Без замены – выборка из перемешанного набора данных. Со заменой – пользователь может указать
num_samplesдля извлечения.- Параметры
-
class torch.utils.data.SubsetRandomSampler(indices, generator=None)[source] -
Случайная выборка элементов из заданного списка индексов без замены.
- Параметры
-
- indices (последовательность) – последовательность индексов
- generator (Generator) – Генератор, используемый для выборки.
-
class torch.utils.data.WeightedRandomSampler(weights, num_samples, replacement=True, generator=None)[source] -
Выборка элементов из
[0,..,len(weights)-1]с заданными вероятностями (весами).- Параметры
-
- weights (последовательность) – последовательность весов, не обязательно суммирующаяся до единицы
- num_samples (int) – количество элементов для выборки
-
replacement (bool) – если
True, элементы выбираются с заменой. В противном случае – без замены, что означает, что когда индекс элемента выбирается для строки, он не может быть выбран снова для этой строки. - generator (Generator) – Генератор, используемый для выборки.
Пример
>>> list(WeightedRandomSampler([0.1, 0.9, 0.4, 0.7, 3.0, 0.6], 5, replacement=True)) [4, 4, 1, 4, 5] >>> list(WeightedRandomSampler([0.9, 0.4, 0.05, 0.2, 0.3, 0.1], 5, replacement=False)) [0, 1, 4, 3, 2]
-
class torch.utils.data.BatchSampler(sampler, batch_size, drop_last)[source] -
Оборачивает другой сортировщик для получения мини-пакета индексов.
- Параметры
Пример
>>> list(BatchSampler(SequentialSampler(range(10)), batch_size=3, drop_last=False)) [[0, 1, 2], [3, 4, 5], [6, 7, 8], [9]] >>> list(BatchSampler(SequentialSampler(range(10)), batch_size=3, drop_last=True)) [[0, 1, 2], [3, 4, 5], [6, 7, 8]]
-
class torch.utils.data.distributed.DistributedSampler(dataset, num_replicas=None, rank=None, shuffle=True, seed=0, drop_last=False)[source] -
Сортировщик, ограничивающий загрузку данных подмножеством набора данных.
Он особенно полезен в сочетании с
torch.nn.parallel.DistributedDataParallel. В таком случае каждый процесс может передать экземплярDistributedSamplerв качестве сортировщикаDataLoaderи загрузить подмножество исходного набора данных, эксклюзивное для него.Примечание
Предполагается, что набор данных имеет постоянный размер и что любой экземпляр его всегда возвращает одни и те же элементы в одном и том же порядке.
- Параметры
-
- dataset – Набор данных, используемый для выборки.
-
num_replicas (int, необязательно) – Количество процессов, участвующих в распределенном обучении. По умолчанию,
world_sizeизвлекается из текущей распределённой группы. -
rank (int, необязательно) – Ранг текущего процесса в
num_replicas. По умолчанию,rankизвлекается из текущей распределённой группы. -
shuffle (bool, необязательно) – Если
True(по умолчанию), сортировщик перемешивает индексы. -
seed (int, необязательно) – случайное семя, используемое для перемешивания сортировщика, если
shuffle=True. Это число должно быть одинаковым во всех процессах в распределённой группе. По умолчанию:0. -
drop_last (bool, необязательно) – если
True, сортировщик отбросит хвост данных, чтобы сделать его равномерно делимым на количество реплик. ЕслиFalse, сортировщик добавит дополнительные индексы, чтобы сделать данные равномерно делимыми на количество реплик. По умолчанию:False.
Предупреждение
В распределённом режиме вызов метода
set_epoch()в начале каждой эпохи до создания итератораDataLoaderнеобходим для правильной работы перемешивания на нескольких эпохах. В противном случае, будет всегда использоваться один и тот же порядок.Пример:
>>> sampler = DistributedSampler(dataset) if is_distributed else None >>> loader = DataLoader(dataset, shuffle=(sampler is None), ... sampler=sampler) >>> for epoch in range(start_epoch, n_epochs): ... if is_distributed: ... sampler.set_epoch(epoch) ... train(loader)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/data.html