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

Класс Dataset

Класс Dataset из модуля torch.utils.data является абстрактным базовым классом для представления набора данных. Он позволяет организовать произвольные данные в виде структуры, доступной для итерации и индексации. Чтобы создать собственный набор данных, необходимо унаследоваться от Dataset и переопределить два метода: __len__, который возвращает размер набора данных, и __getitem__, который возвращает элемент по индексу. Эти методы являются основой для загрузки данных в PyTorch и широко используются вместе с DataLoader.

Синтаксис

class MyDataset(Dataset): def __init__(self, ...): # Инициализация данных pass def __len__(self): # Возвращает общее количество элементов pass def __getitem__(self, idx): # Возвращает элемент по индексу idx pass

Пример

Создадим простой набор данных, который возвращает числа от 0 до 9:

import torch from torch.utils.data import Dataset class NumberDataset(Dataset): def __init__(self, size=10): self.size = size def __len__(self): return self.size def __getitem__(self, idx): return torch.tensor(idx) dataset = NumberDataset() print(len(dataset)) print(dataset[0]) print(dataset[5])

Результат выполнения кода:

10 tensor(0) tensor(5)

Пример

Создадим набор данных, состоящий из пар (признак, метка), где признаки - это случайные векторы, а метки - их сумма:

import torch from torch.utils.data import Dataset class RandomPairDataset(Dataset): def __init__(self, num_samples=5, dim=3): torch.manual_seed(0) self.num_samples = num_samples self.dim = dim self.features = torch.randn(num_samples, dim) self.labels = self.features.sum(dim=1) def __len__(self): return self.num_samples def __getitem__(self, idx): return self.features[idx], self.labels[idx] dataset = RandomPairDataset() print(len(dataset)) feature, label = dataset[0] print(feature) print(label)

Результат выполнения кода:

5 tensor([ 1.5410, -0.2934, -2.1788]) tensor(-0.9312)

Пример

Используем созданный набор данных вместе с DataLoader для организации пакетной загрузки:

import torch from torch.utils.data import Dataset, DataLoader class NumberDataset(Dataset): def __init__(self, size=10): self.size = size def __len__(self): return self.size def __getitem__(self, idx): return torch.tensor(idx) dataset = NumberDataset() dataloader = DataLoader(dataset, batch_size=4, shuffle=True) for batch in dataloader: print(batch)

Результат выполнения кода (порядок может меняться из-за shuffle):

tensor([7, 1, 0, 6]) tensor([4, 9, 5, 2]) tensor([3, 8])

Смотрите также

  • класс Dataset,
    который является основой для создания пользовательских наборов данных
  • метод __getitem__,
    который позволяет получать элемент по индексу
  • метод __len__,
    который возвращает размер набора данных
  • метод __add__,
    который объединяет два набора данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить