Класс Subset
Класс Subset создает подвыборку из существующего датасета на основе переданных индексов. Первым параметром конструктор принимает исходный датасет, вторым - последовательность индексов (список, кортеж, массив NumPy или тензор), которые определяют, какие элементы войдут в подвыборку. Полученный объект ведет себя как обычный датасет: поддерживает индексацию и имеет метод __len__.
Синтаксис
torch.utils.data.Subset(dataset, indices)
Параметры
dataset - исходный датасет, из которого создается подвыборка.
indices - последовательность индексов, определяющих элементы подвыборки. Может быть списком, кортежем, массивом NumPy или тензором.
Пример
Давайте создадим подвыборку из первых пяти элементов существующего датасета:
import torch
from torch.utils.data import Subset, TensorDataset
data = torch.tensor([10, 20, 30, 40, 50, 60, 70, 80, 90, 100])
dataset = TensorDataset(data)
indices = [0, 1, 2, 3, 4]
subset = Subset(dataset, indices)
for i in range(len(subset)):
print(subset[i])
Результат выполнения кода:
(tensor(10),)
(tensor(20),)
(tensor(30),)
(tensor(40),)
(tensor(50),)
Пример
Создадим подвыборку с использованием тензора индексов и перемешанными данными:
import torch
from torch.utils.data import Subset, TensorDataset
torch.manual_seed(0)
data = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12],
[13, 14, 15]
])
dataset = TensorDataset(data)
indices = torch.tensor([2, 0, 4, 1, 3])
subset = Subset(dataset, indices)
for i in range(len(subset)):
print(subset[i])
Результат выполнения кода:
(tensor([7, 8, 9]),)
(tensor([1, 2, 3]),)
(tensor([13, 14, 15]),)
(tensor([4, 5, 6]),)
(tensor([10, 11, 12]),)
Пример
Используем массив NumPy для создания подвыборки из четных индексов:
import torch
import numpy as np
from torch.utils.data import Subset, TensorDataset
data = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
dataset = TensorDataset(data)
indices = np.array([1, 3, 5, 7, 9])
subset = Subset(dataset, indices)
for i in range(len(subset)):
print(subset[i])
Результат выполнения кода:
(tensor(2),)
(tensor(4),)
(tensor(6),)
(tensor(8),)
(tensor(10),)
Смотрите также
-
функцию
random_split,
которая автоматически разбивает датасет на подвыборки заданных размеров -
класс
TensorDataset,
который оборачивает тензоры в датасет для удобной работы -
класс
ConcatDataset,
который объединяет несколько датасетов в один -
класс
Sampler,
который управляет порядком выборки элементов из датасета