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

Метод __len__

Метод __len__ класса DataLoader возвращает общее количество батчей, которое загрузчик данных сгенерирует за одну эпоху. Это число вычисляется на основе длины датасета, размера батча и параметров drop_last и sampler. Метод часто используется для итерации по данным, планирования обучения и отображения прогресса.

Синтаксис

len(dataloader)

Пример

Создадим простой датасет и загрузчик данных. Посмотрим, сколько батчей будет возвращено за одну эпоху при стандартных настройках:

import torch from torch.utils.data import DataLoader, TensorDataset # Create dataset with 100 samples data = torch.randn(100, 10) labels = torch.randint(0, 2, (100,)) dataset = TensorDataset(data, labels) # Create dataloader with batch size 32 dataloader = DataLoader(dataset, batch_size=32) # Get number of batches per epoch num_batches = len(dataloader) print(num_batches)

Результат выполнения кода:

4

В примере датасет из 100 элементов разбивается на батчи по 32 элемента. Получается 3 полных батча (96 элементов) и 1 неполный (4 элемента). Итого len возвращает 4.

Пример

Рассмотрим влияние параметра drop_last на возвращаемое значение. При включении этого параметра последний неполный батч отбрасывается:

import torch from torch.utils.data import DataLoader, TensorDataset # Create dataset with 100 samples data = torch.randn(100, 10) labels = torch.randint(0, 2, (100,)) dataset = TensorDataset(data, labels) # Create dataloader with batch size 32 and drop_last=True dataloader = DataLoader(dataset, batch_size=32, drop_last=True) # Get number of batches per epoch num_batches = len(dataloader) print(num_batches)

Результат выполнения кода:

3

В этом случае неполный батч с 4 элементами отбрасывается, и за эпоху гененируется только 3 полных батча.

Пример

Если размер батча превышает длину датасета, то метод возвращает 1, если drop_last равно False, или 0, если drop_last равно True:

import torch from torch.utils.data import DataLoader, TensorDataset # Create dataset with 5 samples data = torch.randn(5, 10) labels = torch.randint(0, 2, (5,)) dataset = TensorDataset(data, labels) # Create dataloader with batch size 10 dataloader = DataLoader(dataset, batch_size=10) # Get number of batches per epoch num_batches = len(dataloader) print(num_batches) # With drop_last=True dataloader_drop = DataLoader(dataset, batch_size=10, drop_last=True) num_batches_drop = len(dataloader_drop) print(num_batches_drop)

Результат выполнения кода:

1 0

Смотрите также

  • класс DataLoader,
    который создает итератор для загрузки данных
  • атрибут dataset,
    который хранит ссылку на используемый датасет
  • атрибут batch_size,
    который определяет размер батча
  • атрибут drop_last,
    который управляет отбрасыванием последнего неполного батча
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить