Класс Sampler
Класс Sampler является базовым абстрактным классом для всех семплеров в PyTorch. Он определяет, в каком порядке и какие индексы будут подаваться в загрузчик данных. При создании собственного семплера необходимо переопределить метод __iter__, возвращающий итератор по индексам, и метод __len__, возвращающий длину семплера.
Синтаксис
class CustomSampler(Sampler):
def __init__(self, data_source):
self.data_source = data_source
def __iter__(self):
# возвращает итератор по индексам
pass
def __len__(self):
# возвращает количество элементов
pass
Пример
Создадим простой семплер, который возвращает индексы в случайном порядке с возможностью указать зерно для воспроизводимости:
import torch
from torch.utils.data import Sampler, Dataset
import numpy as np
class ReproducibleRandomSampler(Sampler):
def __init__(self, data_source, seed=42):
self.data_source = data_source
self.seed = seed
def __iter__(self):
np.random.seed(self.seed)
indices = np.random.permutation(len(self.data_source))
return iter(indices.tolist())
def __len__(self):
return len(self.data_source)
# создаём простой датасет для демонстрации
class DummyDataset(Dataset):
def __len__(self):
return 10
def __getitem__(self, idx):
return idx
dataset = DummyDataset()
sampler = ReproducibleRandomSampler(dataset, seed=0)
print(list(sampler))
Результат выполнения кода:
[2, 8, 4, 9, 1, 6, 7, 3, 0, 5]
Пример
Реализуем семплер, который возвращает индексы только чётных элементов датасета:
import torch
from torch.utils.data import Sampler, Dataset
class EvenIndicesSampler(Sampler):
def __init__(self, data_source):
self.data_source = data_source
self.even_indices = [i for i in range(len(data_source)) if i % 2 == 0]
def __iter__(self):
return iter(self.even_indices)
def __len__(self):
return len(self.even_indices)
# создаём простой датасет для демонстрации
class DummyDataset(Dataset):
def __len__(self):
return 10
def __getitem__(self, idx):
return idx
dataset = DummyDataset()
sampler = EvenIndicesSampler(dataset)
print(list(sampler))
Результат выполнения кода:
[0, 2, 4, 6, 8]
Пример
Создадим семплер, который возвращает индексы, перемешанные с учётом весов каждого элемента (чем больше вес, тем чаще элемент появляется):
import torch
from torch.utils.data import Sampler, Dataset
import random
class WeightedSampler(Sampler):
def __init__(self, weights, num_samples=None):
self.weights = weights
self.num_samples = num_samples if num_samples else len(weights)
def __iter__(self):
indices = random.choices(
range(len(self.weights)),
weights=self.weights,
k=self.num_samples
)
return iter(indices)
def __len__(self):
return self.num_samples
dataset = DummyDataset()
weights = [0.5, 0.1, 0.4, 0.3, 0.2]
sampler = WeightedSampler(weights, num_samples=10)
print(list(sampler))
Результат выполнения кода (может отличаться при разных запусках):
[0, 2, 0, 3, 0, 2, 4, 0, 3, 2]
Пример
Используем базовый класс Sampler для создания семплера, который возвращает индексы в порядке возрастания, но пропускает каждый второй элемент:
import torch
from torch.utils.data import Sampler, Dataset
class StepSampler(Sampler):
def __init__(self, data_source, step=2):
self.data_source = data_source
self.step = step
def __iter__(self):
indices = list(range(0, len(self.data_source), self.step))
return iter(indices)
def __len__(self):
return len(range(0, len(self.data_source), self.step))
# создаём простой датасет для демонстрации
class DummyDataset(Dataset):
def __len__(self):
return 10
def __getitem__(self, idx):
return idx
dataset = DummyDataset()
sampler = StepSampler(dataset, step=3)
print(list(sampler))
Результат выполнения кода:
[0, 3, 6, 9]
Смотрите также
-
класс
SequentialSampler,
который возвращает индексы в последовательном порядке -
класс
RandomSampler,
который возвращает индексы в случайном порядке -
класс
SubsetRandomSampler,
который создаёт случайную выборку из подмножества индексов -
класс
BatchSampler,
который группирует индексы в батчи заданного размера