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,
который определяет, нужно ли отбрасывать последний неполный пакет