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

batch_sampler

Атрибут batch_sampler класса DataLoader определяет, как именно формируются пакеты (батчи) из индексов сэмплов. Этот атрибут является экземпляром класса-наследника Sampler, который возвращает итератор по спискам индексов для каждого батча. Если вы передаете batch_sampler, параметры batch_size, shuffle, sampler и drop_last должны быть установлены в значение None, так как они определяются самим сэмплером пакетов.

Синтаксис

from torch.utils.data import DataLoader dataloader = DataLoader( dataset, batch_sampler=sampler_object, # другие параметры не должны конфликтовать с batch_sampler )

Пример

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

import torch from torch.utils.data import Dataset, DataLoader, Sampler class MyDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] class MyBatchSampler(Sampler): def __init__(self, data_source, batch_size): self.data_source = data_source self.batch_size = batch_size def __iter__(self): indices = list(range(len(self.data_source))) for i in range(0, len(indices), self.batch_size): yield indices[i:i + self.batch_size] def __len__(self): return (len(self.data_source) + self.batch_size - 1) // self.batch_size dataset = MyDataset([1, 2, 3, 4, 5, 6, 7, 8, 9]) batch_sampler = MyBatchSampler(dataset, batch_size=4) dataloader = DataLoader(dataset, batch_sampler=batch_sampler) for batch in dataloader: print(batch)

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

tensor([1, 2, 3, 4]) tensor([5, 6, 7, 8]) tensor([9])

Пример

Использование встроенного сэмплера пакетов BatchSampler, который комбинирует сэмплер и параметры пакетирования:

import torch from torch.utils.data import Dataset, DataLoader, BatchSampler, SequentialSampler class MyDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset = MyDataset([10, 20, 30, 40, 50, 60, 70]) sampler = SequentialSampler(dataset) batch_sampler = BatchSampler(sampler, batch_size=3, drop_last=False) dataloader = DataLoader(dataset, batch_sampler=batch_sampler) for batch in dataloader: print(batch)

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

tensor([10, 20, 30]) tensor([40, 50, 60]) tensor([70])

Пример

Пример с изменением порядка сэмплов с помощью shuffle:

import torch from torch.utils.data import Dataset, DataLoader, BatchSampler, RandomSampler torch.manual_seed(0) class MyDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset = MyDataset([1, 2, 3, 4, 5, 6, 7, 8]) sampler = RandomSampler(dataset) batch_sampler = BatchSampler(sampler, batch_size=3, drop_last=True) dataloader = DataLoader(dataset, batch_sampler=batch_sampler) for batch in dataloader: print(batch)

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

tensor([5, 4, 8]) tensor([2, 1, 6])

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

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