Класс 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__,
который объединяет два набора данных