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

Класс DataLoader

Класс DataLoader служит для организации загрузки данных из датасета в процессе обучения моделей. Первым параметром он принимает объект датасета (наследник Dataset). Вторым параметром передаётся размер батча batch_size. Также можно настроить количество рабочих процессов num_workers, перемешивание shuffle и другие параметры.

Синтаксис

torch.utils.data.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, multiprocessing_context=None, generator=None, *, prefetch_factor=2, persistent_workers=False, pin_memory_device="" )

Пример

Создадим простой датасет из чисел и загрузим его с помощью DataLoader батчами по 3 элемента:

import torch from torch.utils.data import Dataset, DataLoader class NumberDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] data = [1, 2, 3, 4, 5, 6, 7, 8] dataset = NumberDataset(data) dataloader = DataLoader(dataset, batch_size=3) for batch in dataloader: print(batch)

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

tensor([1, 2, 3]) tensor([4, 5, 6]) tensor([7, 8])

Пример

Используем параметр shuffle для перемешивания данных перед каждой эпохой и параметр drop_last для отбрасывания последнего неполного батча:

import torch from torch.utils.data import Dataset, DataLoader class NumberDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] data = [1, 2, 3, 4, 5, 6, 7] dataset = NumberDataset(data) torch.manual_seed(0) dataloader = DataLoader( dataset, batch_size=3, shuffle=True, drop_last=True ) for batch in dataloader: print(batch)

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

tensor([5, 4, 6]) tensor([2, 3, 1])

Пример

Используем параметр num_workers для параллельной загрузки данных в несколько процессов. Также рассмотрим метод __len__, который возвращает количество батчей в загрузчике:

import torch from torch.utils.data import Dataset, DataLoader class NumberDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] data = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10] dataset = NumberDataset(data) dataloader = DataLoader( dataset, batch_size=4, num_workers=2, shuffle=False ) print(len(dataloader)) for batch in dataloader: print(batch)

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

3 tensor([1, 2, 3, 4]) tensor([5, 6, 7, 8]) tensor([9, 10])

Пример

Рассмотрим работу с атрибутом dataset, который возвращает исходный датасет, и атрибутом batch_size, содержащий размер батча:

import torch from torch.utils.data import Dataset, DataLoader class NumberDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] data = [1, 2, 3, 4, 5] dataset = NumberDataset(data) dataloader = DataLoader(dataset, batch_size=2) print(dataloader.dataset.data) print(dataloader.batch_size)

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

[1, 2, 3, 4, 5] 2

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

  • атрибут dataset,
    который хранит исходный датасет загрузчика
  • атрибут batch_size,
    который определяет размер батча для загрузки
  • метод __iter__,
    который возвращает итератор по батчам данных
  • метод __len__,
    который возвращает количество батчей в загрузчике
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить