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

Класс 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,
    который группирует индексы в батчи заданного размера
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить