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

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