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

Класс BatchSampler

Класс BatchSampler предназначен для формирования пакетов индексов из данных, предоставляемых базовым сэмплером. Первым параметром конструктор принимает объект сэмплера (например, SequentialSampler или RandomSampler). Вторым параметром передаётся размер пакета. Третьим параметром можно указать, нужно ли отбрасывать последний неполный пакет.

Синтаксис

torch.utils.data.BatchSampler( sampler, batch_size, drop_last=False )

Пример

Давайте создадим пакетный сэмплер для последовательной выборки индексов из диапазона 10 элементов размером пакета 3:

import torch from torch.utils.data import BatchSampler from torch.utils.data import SequentialSampler sampler = SequentialSampler(range(10)) batch_sampler = BatchSampler( sampler, batch_size=3, drop_last=False ) for batch in batch_sampler: print(batch)

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

[0, 1, 2] [3, 4, 5] [6, 7, 8] [9]

Пример

Давайте используем BatchSampler с отбрасыванием последнего неполного пакета:

import torch from torch.utils.data import BatchSampler from torch.utils.data import SequentialSampler sampler = SequentialSampler(range(10)) batch_sampler = BatchSampler( sampler, batch_size=3, drop_last=True ) for batch in batch_sampler: print(batch)

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

[0, 1, 2] [3, 4, 5] [6, 7, 8]

Пример

Давайте используем BatchSampler с перемешиванием данных с помощью RandomSampler:

import torch from torch.utils.data import BatchSampler from torch.utils.data import RandomSampler torch.manual_seed(0) sampler = RandomSampler(range(10)) batch_sampler = BatchSampler( sampler, batch_size=3, drop_last=False ) for batch in batch_sampler: print(batch)

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

[6, 9, 3] [4, 0, 5] [2, 7, 1] [8]

Пример

Давайте используем BatchSampler вместе с загрузчиком данных DataLoader для загрузки мини-пакетов:

import torch from torch.utils.data import DataLoader from torch.utils.data import TensorDataset from torch.utils.data import BatchSampler from torch.utils.data import SequentialSampler data = torch.tensor([ [1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15], [16, 17, 18], ]) dataset = TensorDataset(data) sampler = SequentialSampler(dataset) batch_sampler = BatchSampler( sampler, batch_size=2, drop_last=True ) dataloader = DataLoader( dataset, batch_sampler=batch_sampler ) for batch in dataloader: print(batch)

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

[tensor([ [1, 2, 3], [4, 5, 6], ])] [tensor([ [7, 8, 9], [10, 11, 12], ])] [tensor([ [13, 14, 15], [16, 17, 18], ])]

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

  • класс Sampler,
    который является базовым классом для всех сэмплеров
  • класс SequentialSampler,
    который возвращает индексы в последовательном порядке
  • класс RandomSampler,
    который возвращает индексы в случайном порядке
  • класс SubsetRandomSampler,
    который возвращает случайные индексы из подмножества данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить