Атрибут drop_last
Атрибут drop_last класса DataLoader определяет,
будет ли отброшен последний батч, если размер набора данных
не делится нацело на размер батча batch_size.
Если drop_last=True, то последний неполный батч
не включается в итерацию. Если drop_last=False
(значение по умолчанию), то последний батч возвращается,
даже если он содержит меньше элементов, чем batch_size.
Этот параметр полезен при обучении моделей, когда требуется
одинаковый размер батча на каждой итерации.
Синтаксис
torch.utils.data.DataLoader(
dataset,
batch_size=1,
drop_last=False,
...
)
Пример
Создадим простой набор данных из 10 элементов и загрузчик с batch_size=3 и drop_last=False (по умолчанию):
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.arange(10)
dataset = TensorDataset(data)
dataloader = DataLoader(
dataset,
batch_size=3,
drop_last=False
)
for batch in dataloader:
print(batch[0])
Результат выполнения кода:
tensor([0, 1, 2])
tensor([3, 4, 5])
tensor([6, 7, 8])
tensor([9])
Последний батч содержит только один элемент, так как 10 не делится на 3 нацело.
Пример
Теперь установим drop_last=True, чтобы отбросить последний неполный батч:
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.arange(10)
dataset = TensorDataset(data)
dataloader = DataLoader(
dataset,
batch_size=3,
drop_last=True
)
for batch in dataloader:
print(batch[0])
Результат выполнения кода:
tensor([0, 1, 2])
tensor([3, 4, 5])
tensor([6, 7, 8])
Последний батч с одним элементом был отброшен, и загрузчик вернул только 3 полных батча.
Пример
Использование drop_last совместно с shuffle=True. Перемешивание данных может влиять на то, какие элементы попадут в последний неполный батч:
import torch
from torch.utils.data import DataLoader, TensorDataset
torch.manual_seed(0)
data = torch.arange(10)
dataset = TensorDataset(data)
dataloader = DataLoader(
dataset,
batch_size=4,
shuffle=True,
drop_last=True
)
for batch in dataloader:
print(batch[0])
Результат выполнения кода:
tensor([6, 8, 7, 9])
tensor([4, 2, 1, 3])
При drop_last=True мы получили только два полных батча по 4 элемента. Элемент 0 и 5 попали в отброшенный неполный батч.
Пример
Важно учитывать, что длина загрузчика (количество батчей) зависит от значения drop_last. Проверим это с помощью функции len:
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.arange(10)
dataset = TensorDataset(data)
dataloader_false = DataLoader(
dataset,
batch_size=3,
drop_last=False
)
print(len(dataloader_false))
dataloader_true = DataLoader(
dataset,
batch_size=3,
drop_last=True
)
print(len(dataloader_true))
Результат выполнения кода:
4
3
В первом случае длина равна 4 (три полных батча и один неполный), во втором - 3 (только полные батчи).
Смотрите также
-
класс
DataLoader,
который реализует загрузку данных по батчам -
атрибут
batch_size,
который определяет размер батча -
атрибут
dataset,
который содержит сами данные -
метод
__iter__,
который возвращает итератор по батчам