Метод __add__ класса Dataset
Метод __add__ класса Dataset позволяет объединять два датасета в один. В результате работы метода возвращается новый объект типа ConcatDataset, который содержит элементы обоих исходных датасетов. Метод вызывается при использовании оператора сложения + между двумя датасетами.
Синтаксис
dataset = dataset1 + dataset2
Пример
Давайте создадим два простых датасета и объединим их с помощью метода __add__:
import torch
from torch.utils.data import Dataset
class SimpleDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return torch.tensor(self.data[idx])
dataset1 = SimpleDataset([1, 2, 3, 4, 5])
dataset2 = SimpleDataset([6, 7, 8, 9, 10])
combined = dataset1 + dataset2
print(len(combined))
Результат выполнения кода:
10
Пример
Давайте проверим, что элементы объединённого датасета соответствуют элементам исходных датасетов:
import torch
from torch.utils.data import Dataset
class SimpleDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return torch.tensor(self.data[idx])
dataset1 = SimpleDataset([1, 2, 3])
dataset2 = SimpleDataset([4, 5, 6])
combined = dataset1 + dataset2
for i in range(len(combined)):
print(combined[i])
Результат выполнения кода:
tensor(1)
tensor(2)
tensor(3)
tensor(4)
tensor(5)
tensor(6)
Пример
Давайте объединим три датасета с помощью цепочки операций сложения:
import torch
from torch.utils.data import Dataset
class SimpleDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return torch.tensor(self.data[idx])
dataset1 = SimpleDataset([1, 2])
dataset2 = SimpleDataset([3, 4])
dataset3 = SimpleDataset([5, 6])
combined = dataset1 + dataset2 + dataset3
print(len(combined))
print(combined[0])
print(combined[3])
print(combined[5])
Результат выполнения кода:
6
tensor(1)
tensor(4)
tensor(6)
Смотрите также
-
класс
Dataset,
который является базовым классом для создания пользовательских датасетов -
метод
__getitem__,
который возвращает элемент датасета по индексу -
метод
__len__,
который возвращает размер датасета -
класс
ConcatDataset,
который используется для объединения нескольких датасетов