torch.utils.data
В основе PyTorch утилиты загрузки данных лежит класс torch.utils.data.DataLoader. Он представляет собой итерируемый объект над набором данных, поддерживающий
- стилизированные и итерируемые наборы данных,
- настраиваемый порядок загрузки данных,
- автоматическое создание батчей,
- загрузку данных с одним и несколькими процессами,
- автоматическое закрепление данных в памяти.
Эти параметры конфигурируются аргументами конструктора 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 поддерживает два разных типа наборов данных:
Стилизированные наборы данных
Стилизированный набор данных — это набор данных, реализующий протоколы __getitem__() и __len__(), и представляющий собой отображение (возможно, не целочисленных) индексов/ключей на образцы данных.
Например, такой набор данных при обращении с помощью dataset[idx], может читать idx-й образ и его соответствующую метку из папки на диске.
См. Dataset для получения более подробной информации.
Итерируемые наборы данных
Итерируемый набор данных — это экземпляр подкласса IterableDataset, который реализует протокол __iter__(), и представляет собой итерируемый объект над образцами данных. Этот тип наборов данных особенно подходит для случаев, когда случайный доступ к данным является дорогостоящим или даже маловероятным, и где размер пакета зависит от полученных данных.
Например, такой набор данных при вызове iter(dataset), может возвращать поток данных, считывая их из базы данных, удаленного сервера или даже журналов, генерируемых в реальном времени.
См. IterableDataset для получения более подробной информации.
Примечание
При использовании IterableDataset с загрузкой данных с несколькими процессами. Один и тот же объект набора данных дублируется в каждом процессе, и поэтому дубликаты должны быть сконфигурированы по-разному, чтобы избежать дублирования данных. Смотрите документацию IterableDataset для получения информации о том, как этого достичь.
Порядок загрузки данных и выборка
Для итерируемых наборов данных порядок загрузки данных полностью контролируется пользователем. Это позволяет проще реализовывать чтение кусками и динамический размер пакета (например, путём выдачи батчевого образца в каждый момент).
Остальная часть этого раздела касается стилизированных наборов данных. Классы torch.utils.data.Sampler используются для указания последовательности индексов/ключей, используемых при загрузке данных. Они представляют собой итерируемые объекты по индексам наборов данных. Например, в случае со случайным градиентным спуском (SGD), Sampler может случайным образом переставлять список индексов и выдавать каждый по одному, или выдавать небольшое их количество для мини-батчевого SGD.
Последовательный или перемешанный выборщик будет автоматически создан на основе аргумента shuffle к DataLoader. В качестве альтернативы, пользователи могут использовать аргумент sampler для указания настраиваемого объекта Sampler, который каждый раз выдаёт следующий индекс/ключ для извлечения.
Настраиваемый Sampler, который выдаёт список индексов пакета за раз, может быть передан в качестве аргумента batch_sampler. Автоматическое создание батчей также можно включить с помощью аргументов batch_size и drop_last. Подробнее об этом см. в следующем разделе .
Примечание
Ни sampler ни batch_sampler не совместимы с итерируемыми наборами данных, так как такие наборы данных не имеют понятия о ключе или индексе.
Загрузка данных в партиях и без партий
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 отбрасывает последнюю неполную партию каждой копии набора данных в рабочем процессе.
После извлечения списка образцов с помощью индексов из sampler, функция, переданная в качестве аргумента 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.
Загрузка данных с использованием одного или нескольких процессов
Класс 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 для индивидуальной настройки каждой реплики набора данных и определения того, выполняется ли код в процессе-рабочем. Например, это особенно полезно при фрагментации набора данных.
Для наборов данных типа map основной процесс генерирует индексы с помощью sampler и отправляет их рабочим процессам. Таким образом, любая случайная перетасовка выполняется в основном процессе, который управляет загрузкой, назначая индексы для загрузки.
Для наборов данных типа iterable, поскольку каждый рабочий процесс получает реплику объекта dataset, простая загрузка с несколькими процессами часто приводит к дублированию данных. Используя torch.utils.data.get_worker_info() и/или worker_init_fn, пользователи могут настраивать каждую реплику независимо. (См. документацию IterableDataset о том, как этого добиться). По тем же причинам в загрузке с несколькими процессами аргумент drop_last отбрасывает последнюю неполную партию каждой реплики набора данных типа iterable в рабочем процессе.
Рабочие процессы завершаются по достижении конца итерации или при удалении итератора.
Предупреждение
В целом не рекомендуется возвращать тензоры 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() и использовать его для задания начального значения генератора случайных чисел других библиотек перед загрузкой данных.
Фиксирование памяти
Копии с хоста на GPU значительно быстрее, когда они происходят из фиксированной (заблокированной в памяти) памяти. См. Использование буферов с фиксированной памятью для получения более подробной информации о том, когда и как использовать фиксированную память в целом.
Для загрузки данных передача pin_memory=True в DataLoader автоматически поместит полученные тензоры данных в фиксированную память, что позволит ускорить передачу данных на CUDA-совместимые GPU.
Логика по умолчанию по фиксированию памяти распознает только тензоры и отображения и итерируемые объекты, содержащие тензоры. По умолчанию, если логика фиксирования видит пакет пользовательского типа (что произойдет, если у вас есть 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=2, persistent_workers=False, pin_memory_device='')[source] -
Загрузчик данных. Объединяет набор данных и выборщик и предоставляет итерируемый объект над заданным набором данных.
DataLoaderподдерживает наборы данных как в стиле отображения, так и в стиле итерации с однопроцессной или многопроцессной загрузкой, настройкой порядка загрузки и дополнительной автоматической пакетной загрузкой (сопоставлением) и фиксированием памяти.См. страницу документации
torch.utils.dataдля получения дополнительной информации.- Параметры:
-
- dataset (Dataset) – набор данных, из которого загружать данные.
-
batch_size (int, необязательно) – количество образцов в пакете для загрузки (по умолчанию:
1). -
shuffle (bool, необязательно) – устанавливается в
Trueдля повторной перетасовки данных на каждой эпохе (по умолчанию:False). -
sampler (Sampler или Iterable, необязательно) – определяет стратегию выборки образцов из набора данных. Может быть любым
Iterableс реализованным__len__. Если указано,shuffleне должно быть указано. -
batch_sampler (Sampler или Iterable, необязательно) – как
sampler, но возвращает пакет индексов за раз. Взаимоисключающее сbatch_size,shuffle,sampler, иdrop_last. -
num_workers (int, необязательно) – количество дочерних процессов для загрузки данных.
0означает, что данные будут загружаться в основном процессе. (по умолчанию:0). - collate_fn (Callable, необязательно) – объединяет список образцов, чтобы сформировать мини-пакет тензора(ов). Используется при использовании пакетной загрузки из набора данных в стиле отображения.
-
pin_memory (bool, необязательно) – Если
True, загрузчик данных скопирует тензоры в память устройства/CUDA перед возвратом. Если ваши элементы данных — пользовательского типа, или вашcollate_fnвозвращает пакет пользовательского типа, см. пример ниже. -
drop_last (bool, необязательно) – установить в
Trueдля пропуска последнего неполного пакета, если размер набора данных не делится на размер пакета. ЕслиFalseи размер набора данных не делится на размер пакета, тогда последний пакет будет меньше. (по умолчанию:False). -
timeout (числовое, необязательно) – если положительное, значение таймаута для сбора пакета от рабочих процессов. Должно быть всегда неотрицательным. (по умолчанию:
0). -
worker_init_fn (Callable, необязательно) – Если не
None, это будет вызываться в каждом дочернем процессе рабочего процесса с идентификатором рабочего процесса (целое число в[0, num_workers - 1]) в качестве входных данных после инициализации и до загрузки данных. (по умолчанию:None). - и т.д. (остальные параметры)
Предупреждение
Если используется метод запуска
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.Примечание
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] -
Набор данных, содержащий тензоры.
Каждый образец будет извлекаться путем индексирования тензоров по первому измерению.
- Parameters:
-
*tensors (Tensor) – тензоры, имеющие одинаковый размер первого измерения.
-
class torch.utils.data.ConcatDataset(datasets)[source] -
Набор данных как конкатенация нескольких наборов данных.
Этот класс полезен для объединения различных существующих наборов данных.
- Parameters:
-
datasets (sequence) – Список наборов данных, которые необходимо конкатенировать
-
class torch.utils.data.ChainDataset(datasets)[source] -
Набор данных для цепочки нескольких
IterableDataset.Этот класс полезен для объединения различных существующих потоков наборов данных. Операция цепочки выполняется динамически, поэтому конкатенация масштабируемых наборов данных с этим классом будет эффективной.
- Parameters:
-
datasets (iterable of IterableDataset) – наборы данных, которые необходимо объединить в цепочку
-
class torch.utils.data.Subset(dataset, indices)[source] -
Подмножество набора данных в указанных индексах.
- Parameters:
-
- dataset (Dataset) – Весь набор данных
- indices (sequence) – Индексы в полном наборе, выбранные для подмножества
-
torch.utils.data._utils.collate.collate(batch, *, collate_fn_map=None)[source] -
Общая функция объединения, которая обрабатывает тип элементов в каждом наборе и открывает реестр функций для работы со специфическими типами элементов.
default_collate_fn_mapпредоставляет функции объединения по умолчанию для тензоров, массивов NumPy, чисел и строк.- Parameters:
-
- 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, …]), …]
- Parameters:
-
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(). Смотрите описание там для более подробной информации.- Parameters:
-
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для задания семян других библиотек, используемых в коде набора данных. -
-
torch.utils.data.random_split(dataset, lengths, generator=<torch._C.Generator object>)[source] -
Случайное разделение набора данных на новые непересекающиеся наборы данных заданных длин.
Если задан список дробей, сумма которых равна 1, длины будут вычисляться автоматически как floor(frac * len(dataset)) для каждой предоставленной дроби.
После вычисления длин, если есть остатки, 1 единица будет распределена по длине в порядке циклического обхода, пока не останется остатков.
Можно зафиксировать генератор для воспроизводимых результатов, например:
>>> random_split(range(10), [3, 7], generator=torch.Generator().manual_seed(42)) >>> random_split(range(30), [0.3, 0.3, 0.4], generator=torch.Generator( ... ).manual_seed(42))
-
class torch.utils.data.Sampler(data_source)[source] -
Базовый класс для всех Samplers.
Каждый подкласс Sampler должен предоставить метод
__iter__(), обеспечивающий способ итерирования по индексам элементов набора данных, и метод__len__(), возвращающий длину возвращаемых итераторов.Примечание
Метод
__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] -
Оборачивает другой sampler, чтобы возвращать мини-пакет индексов.
- Параметры:
Пример
>>> 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, optional) – Количество процессов, участвующих в распределённом обучении. По умолчанию извлекается из текущей распределённой группы.
-
rank (int, optional) – Ранг текущего процесса в
num_replicas. По умолчанию извлекается из текущей распределённой группы. -
shuffle (bool, optional) – Если
True(по умолчанию), сэмплер перемешивает индексы. -
seed (int, optional) – случайное семя, используемое для перемешивания сэмплера, если
shuffle=True. Это число должно быть одинаковым во всех процессах распределённой группы. По умолчанию:0. -
drop_last (bool, optional) – если
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/1.13/data.html