Атрибут num_workers
Атрибут num_workers класса DataLoader задает количество подпроцессов,
которые будут использоваться для параллельной загрузки данных.
Чем больше значение, тем быстрее может загружаться данные, но возрастает нагрузка на память и процессор.
Значение по умолчанию 0 означает, что загрузка данных будет происходить в основном процессе.
Синтаксис
torch.utils.data.DataLoader(
dataset,
batch_size=1,
num_workers=0
)
Пример
Давайте создадим даталоадер с одним рабочим процессом:
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]])
targets = torch.tensor([0, 1, 0, 1])
dataset = TensorDataset(data, targets)
dataloader = DataLoader(dataset, batch_size=2, num_workers=1)
print(dataloader.num_workers)
Результат выполнения кода:
1
Пример
Давайте сравним время загрузки данных с разным количеством рабочих процессов:
import torch
from torch.utils.data import DataLoader, TensorDataset
import time
torch.manual_seed(0)
data = torch.randn(10000, 100)
targets = torch.randint(0, 2, (10000,))
dataset = TensorDataset(data, targets)
for num_workers in [0, 2, 4]:
start = time.time()
dataloader = DataLoader(
dataset,
batch_size=64,
num_workers=num_workers
)
for batch in dataloader:
pass
end = time.time()
print(f'num_workers={num_workers}: {end - start:.4f}s')
Результат выполнения кода:
num_workers=0: 0.4523s
num_workers=2: 0.2187s
num_workers=4: 0.1954s
Пример
Давайте создадим даталоадер с большим количеством рабочих процессов для интенсивной загрузки данных:
import torch
from torch.utils.data import DataLoader, Dataset
class CustomDataset(Dataset):
def __init__(self, size):
self.size = size
self.data = torch.randn(size, 3, 224, 224)
self.targets = torch.randint(0, 10, (size,))
def __len__(self):
return self.size
def __getitem__(self, idx):
return self.data[idx], self.targets[idx]
dataset = CustomDataset(1000)
dataloader = DataLoader(
dataset,
batch_size=32,
num_workers=8,
pin_memory=True
)
print(f'num_workers: {dataloader.num_workers}')
print(f'batch_size: {dataloader.batch_size}')
Результат выполнения кода:
num_workers: 8
batch_size: 32
Смотрите также
-
класс
DataLoader,
который управляет загрузкой данных в PyTorch -
атрибут
batch_size,
который определяет размер батча -
атрибут
pin_memory,
который ускоряет передачу данных на GPU -
метод
__iter__,
который возвращает итератор по данным