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

Атрибут 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__,
    который возвращает итератор по данным
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить