Класс 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__,
который возвращает количество батчей в загрузчике