Метод __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,
который управляет отбрасыванием последнего неполного батча