Метод __iter__ класса DataLoader
Метод __iter__ класса DataLoader возвращает объект итератора,
который последовательно выдает батчи данных из переданного датасета.
Этот метод вызывается неявно при использовании цикла for или
функции iter. Он учитывает все параметры загрузчика:
размер батча, способ семплирования, количество рабочих процессов,
перемешивание данных и другие настройки.
Синтаксис
iterator = iter(dataloader)
# или
for batch in dataloader:
# обработка батча
Пример
Создадим простой датасет из чисел и получим итератор с помощью метода __iter__:
import torch
from torch.utils.data import DataLoader, TensorDataset
# Создаем данные
data = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
dataset = TensorDataset(data)
# Создаем загрузчик с батчами по 3 элемента
dataloader = DataLoader(dataset, batch_size=3, shuffle=False)
# Получаем итератор через __iter__
iterator = iter(dataloader)
# Получаем первый батч
batch1 = next(iterator)
print(batch1)
Результат выполнения кода:
[tensor([1, 2, 3])]
Пример
Переберем все батчи с помощью цикла for, который неявно вызывает __iter__:
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.tensor([10, 20, 30, 40, 50, 60, 70])
dataset = TensorDataset(data)
dataloader = DataLoader(dataset, batch_size=2, shuffle=False)
for i, batch in enumerate(dataloader):
print(f"Batch {i}: {batch}")
Результат выполнения кода:
"Batch 0: [tensor([10, 20])]"
"Batch 1: [tensor([30, 40])]"
"Batch 2: [tensor([50, 60])]"
"Batch 3: [tensor([70])]"
Пример
При использовании нескольких рабочих процессов (num_workers) метод __iter__ создает итераторы для каждого процесса:
import torch
from torch.utils.data import DataLoader, TensorDataset
torch.manual_seed(0)
data = torch.arange(0, 20)
dataset = TensorDataset(data)
dataloader = DataLoader(
dataset,
batch_size=4,
shuffle=True,
num_workers=2
)
for batch in dataloader:
print(batch)
break # берем только первый батч
Результат выполнения кода (может отличаться из-за случайного перемешивания):
[tensor([15, 12, 13, 10])]
Смотрите также
-
метод
__len__,
который возвращает количество батчей в загрузчике -
атрибут
dataset,
который хранит ссылку на исходный датасет -
атрибут
batch_size,
который определяет размер батча -
атрибут
num_workers,
который задает количество подпроцессов для загрузки данных