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

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