Класс ConcatDataset
Класс ConcatDataset служит для объединения нескольких наборов данных в один. Это полезно, когда нужно работать с данными, разбитыми на несколько частей, как с единым целым. Первым параметром конструктор принимает последовательность (список или кортеж) объектов наборов данных, которые необходимо объединить.
Синтаксис
from torch.utils.data import ConcatDataset
concat_dataset = ConcatDataset([dataset1, dataset2, ...])
Пример
Давайте создадим два простых набора данных в виде списков и объединим их с помощью ConcatDataset:
import torch
from torch.utils.data import ConcatDataset, DataLoader
# Создаем два списка с данными
data1 = [1, 2, 3, 4, 5]
data2 = [6, 7, 8, 9, 10]
# Объединяем их в один набор данных
concat_dataset = ConcatDataset([data1, data2])
# Создаем загрузчик данных
loader = DataLoader(concat_dataset, batch_size=3, shuffle=False)
# Выводим элементы
for batch in loader:
print(batch)
Результат выполнения кода:
tensor([1, 2, 3])
tensor([4, 5, 6])
tensor([7, 8, 9])
tensor([10])
Пример
Теперь объединим два набора данных типа TensorDataset, содержащих пары признак-метка, и пройдём по ним загрузчиком:
import torch
from torch.utils.data import ConcatDataset, TensorDataset, DataLoader
# Первый набор данных
x1 = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
y1 = torch.tensor([0, 1])
dataset1 = TensorDataset(x1, y1)
# Второй набор данных
x2 = torch.tensor([[5.0, 6.0], [7.0, 8.0]])
y2 = torch.tensor([0, 1])
dataset2 = TensorDataset(x2, y2)
# Объединение
concat_dataset = ConcatDataset([dataset1, dataset2])
# Загрузчик
loader = DataLoader(concat_dataset, batch_size=2, shuffle=True)
torch.manual_seed(0)
# Вывод
for x_batch, y_batch in loader:
print("x:", x_batch)
print("y:", y_batch)
Результат выполнения кода:
x: tensor([[5., 6.],
[1., 2.]])
y: tensor([0, 0])
x: tensor([[3., 4.],
[7., 8.]])
y: tensor([1, 1])
Пример
Проверим длину объединённого набора данных с помощью функции len:
import torch
from torch.utils.data import ConcatDataset
data1 = torch.tensor([1, 2, 3])
data2 = torch.tensor([4, 5, 6, 7])
concat_dataset = ConcatDataset([data1, data2])
print(len(concat_dataset))
Результат выполнения кода:
7
Смотрите также
-
класс
TensorDataset,
который создает набор данных из тензоров -
класс
ChainDataset,
который используется для объединения итерируемых наборов данных -
класс
Subset,
который выделяет подмножество из набора данных -
функцию
random_split,
которая разбивает набор данных на части случайным образом