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

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