Класс SubsetRandomSampler
Класс SubsetRandomSampler создает сэмплер, который
случайным образом выбирает индексы из переданного подмножества.
Первым параметром конструктор принимает список или диапазон индексов,
которые составляют подмножество. Вторым параметром можно передать
генератор случайных чисел для воспроизводимости.
Этот сэмплер используется с DataLoader для получения
перемешанных батчей только из определенной части датасета.
Синтаксис
torch.utils.data.SubsetRandomSampler(indices, generator=None)
Пример
Давайте создадим сэмплер для случайной выборки индексов из списка:
import torch
from torch.utils.data import SubsetRandomSampler
# Создаем сэмплер для индексов 0, 1, 2, 3, 4
sampler = SubsetRandomSampler([0, 1, 2, 3, 4])
# Получаем список индексов (для демонстрации)
indices = list(sampler)
print(indices)
Результат выполнения кода:
[2, 4, 0, 3, 1]
Пример
Давайте используем SubsetRandomSampler с датасетом:
import torch
from torch.utils.data import DataLoader, TensorDataset, SubsetRandomSampler
# Создаем датасет
data = torch.tensor([
[1, 2],
[3, 4],
[5, 6],
[7, 8],
[9, 10],
])
dataset = TensorDataset(data)
# Создаем сэмплер для индексов 0, 1, 2
sampler = SubsetRandomSampler([0, 1, 2])
# Создаем DataLoader с сэмплером
loader = DataLoader(dataset, batch_size=2, sampler=sampler)
# Итерируем по данным
for batch in loader:
print(batch)
Результат выполнения кода:
[tensor([[5, 6],
[1, 2]])]
[tensor([[3, 4]])]
Пример
Давайте используем сэмплер с фиксированным зерном для воспроизводимости:
import torch
from torch.utils.data import DataLoader, TensorDataset, SubsetRandomSampler
torch.manual_seed(0)
# Создаем датасет
data = torch.tensor([
[1, 2],
[3, 4],
[5, 6],
[7, 8],
[9, 10],
])
dataset = TensorDataset(data)
# Создаем сэмплер с генератором
generator = torch.Generator()
generator.manual_seed(42)
sampler = SubsetRandomSampler([0, 1, 2, 3, 4], generator=generator)
# Создаем DataLoader
loader = DataLoader(dataset, batch_size=3, sampler=sampler)
for batch in loader:
print(batch)
Результат выполнения кода:
[tensor([[5, 6],
[9, 10],
[3, 4]])]
[tensor([[7, 8],
[1, 2]])]
Пример
Давайте используем SubsetRandomSampler для валидационной выборки:
import torch
from torch.utils.data import DataLoader, TensorDataset, SubsetRandomSampler
torch.manual_seed(0)
# Создаем датасет
data = torch.tensor([
[1, 2],
[3, 4],
[5, 6],
[7, 8],
[9, 10],
[11, 12],
])
dataset = TensorDataset(data)
# Индексы для валидации (последние 2)
val_indices = [4, 5]
# Создаем сэмплер для валидации
val_sampler = SubsetRandomSampler(val_indices)
# Создаем валидационный DataLoader
val_loader = DataLoader(
dataset,
batch_size=2,
sampler=val_sampler,
)
for batch in val_loader:
print(batch)
Результат выполнения кода:
[tensor([[11, 12],
[9, 10]])]
Смотрите также
-
класс
Sampler,
который является базовым классом для всех сэмплеров -
класс
RandomSampler,
который создает сэмплер для случайной выборки из всего датасета -
класс
SequentialSampler,
который создает сэмплер для последовательной выборки -
функцию
random_split,
которая разбивает датасет на подмножества случайным образом