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

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