Атрибут sampler
Атрибут sampler класса DataLoader определяет,
как именно выбираются индексы элементов из датасета для формирования
каждого батча. Этот атрибут принимает объект, наследующий от
Sampler, и позволяет гибко управлять порядком и
способом выборки данных.
Если sampler не указан явно, DataLoader
автоматически выбирает стратегию в зависимости от значения
параметра shuffle: при shuffle=True
используется RandomSampler, а при
shuffle=False - SequentialSampler.
Атрибут доступен только для чтения и определяется на этапе
инициализации загрузчика.
Синтаксис
from torch.utils.data import DataLoader, Dataset
dataloader = DataLoader(
dataset,
sampler=sampler_object,
...
)
Пример
Создадим простой датасет и загрузчик с явным указанием
SequentialSampler для последовательного обхода
данных:
import torch
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.sampler import SequentialSampler
class SimpleDataset(Dataset):
def __init__(self):
self.data = [1, 2, 3, 4, 5, 6]
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
sampler = SequentialSampler(dataset)
dataloader = DataLoader(dataset, sampler=sampler, batch_size=2)
for batch in dataloader:
print(batch)
Результат выполнения кода:
tensor([1, 2])
tensor([3, 4])
tensor([5, 6])
Пример
Используем RandomSampler с фиксированным зерном
для воспроизводимой случайной выборки:
import torch
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.sampler import RandomSampler
torch.manual_seed(0)
class SimpleDataset(Dataset):
def __init__(self):
self.data = [10, 20, 30, 40, 50, 60]
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
sampler = RandomSampler(dataset)
dataloader = DataLoader(dataset, sampler=sampler, batch_size=3)
for batch in dataloader:
print(batch)
Результат выполнения кода:
tensor([50, 30, 10])
tensor([60, 20, 40])
Пример
Создадим собственный сэмплер, который выбирает элементы в обратном порядке:
import torch
from torch.utils.data import DataLoader, Dataset, Sampler
class ReverseSampler(Sampler):
def __init__(self, data_source):
self.data_source = data_source
def __iter__(self):
return iter(range(len(self.data_source) - 1, -1, -1))
def __len__(self):
return len(self.data_source)
class SimpleDataset(Dataset):
def __init__(self):
self.data = ['a', 'b', 'c', 'd', 'e', 'f']
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
sampler = ReverseSampler(dataset)
dataloader = DataLoader(dataset, sampler=sampler, batch_size=2)
for batch in dataloader:
print(batch)
Результат выполнения кода:
['f', 'e']
['d', 'c']
['b', 'a']
Смотрите также
-
атрибут
batch_sampler,
который генерирует целые батчи индексов -
атрибут
dataset,
который хранит источник данных для загрузчика -
атрибут
batch_size,
который определяет размер батча -
класс
DataLoader,
который инкапсулирует весь процесс загрузки данных