Класс SequentialSampler
Класс SequentialSampler предназначен для создания сэмплера,
который возвращает индексы элементов датасета в последовательном порядке.
Этот сэмплер используется по умолчанию в DataLoader,
если не указан другой сэмплер или параметр shuffle не установлен в True.
При инициализации класс принимает источник данных или его длину.
Синтаксис
torch.utils.data.SequentialSampler(data_source)
Параметры:
-
data_source- датасет или другой итерируемый объект, из которого извлекаются индексы.
Пример
Создадим простой датасет и применим к нему последовательный сэмплер:
import torch
from torch.utils.data import SequentialSampler, TensorDataset
data = torch.tensor([10, 20, 30, 40, 50])
dataset = TensorDataset(data)
sampler = SequentialSampler(dataset)
for idx in sampler:
print(idx, data[idx].item())
Результат выполнения кода:
0 10
1 20
2 30
3 40
4 50
Пример
Используем сэмплер в загрузчике данных для последовательной подачи батчей:
import torch
from torch.utils.data import DataLoader, TensorDataset, SequentialSampler
data = torch.arange(1, 11)
dataset = TensorDataset(data)
sampler = SequentialSampler(dataset)
loader = DataLoader(dataset, batch_size=3, sampler=sampler)
for batch in loader:
print(batch)
Результат выполнения кода:
tensor([
[1],
[2],
[3],
])
tensor([
[4],
[5],
[6],
])
tensor([
[7],
[8],
[9],
])
tensor([
[10],
])
Пример
Сравним работу SequentialSampler и RandomSampler:
import torch
from torch.utils.data import SequentialSampler, RandomSampler, TensorDataset
torch.manual_seed(0)
data = torch.tensor([1, 2, 3, 4, 5])
dataset = TensorDataset(data)
seq_sampler = SequentialSampler(dataset)
rand_sampler = RandomSampler(dataset)
print("SequentialSampler:")
for idx in seq_sampler:
print(idx, end=" ")
print("\nRandomSampler:")
for idx in rand_sampler:
print(idx, end=" ")
Результат выполнения кода:
SequentialSampler:
0 1 2 3 4
RandomSampler:
4 0 1 3 2
Смотрите также
-
класс
Sampler,
базовый класс для всех сэмплеров -
класс
RandomSampler,
сэмплер для случайной выборки индексов -
класс
BatchSampler,
сэмплер для формирования батчей по заданным индексам -
класс
TensorDataset,
датасет для хранения тензоров в виде выборок