Класс TensorDataset
Класс TensorDataset из модуля torch.utils.data позволяет обернуть один или несколько тензоров в объект, совместимый с загрузчиком данных DataLoader. При обращении по индексу он возвращает кортеж, содержащий элементы из каждого переданного тензора на соответствующей позиции. Все переданные тензоры должны иметь одинаковый размер по первому измерению.
Синтаксис
torch.utils.data.TensorDataset(*tensors)
Конструктор принимает произвольное количество тензоров в качестве позиционных аргументов.
Пример
Создадим набор данных из двух тензоров - признаков и меток:
import torch
from torch.utils.data import TensorDataset
features = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
[5.0, 6.0],
])
labels = torch.tensor([0, 1, 0])
dataset = TensorDataset(features, labels)
Пример
Получим первый элемент из созданного набора данных:
import torch
from torch.utils.data import TensorDataset
features = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
[5.0, 6.0],
])
labels = torch.tensor([0, 1, 0])
dataset = TensorDataset(features, labels)
first = dataset[0]
print(first)
Результат выполнения кода:
(tensor([1., 2.]), tensor(0))
Пример
Используем набор данных вместе с загрузчиком DataLoader для итерации по мини-батчам:
import torch
from torch.utils.data import TensorDataset, DataLoader
features = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
[5.0, 6.0],
[7.0, 8.0],
])
labels = torch.tensor([0, 1, 0, 1])
dataset = TensorDataset(features, labels)
loader = DataLoader(dataset, batch_size=2, shuffle=True)
for batch_features, batch_labels in loader:
print(batch_features, batch_labels)
Результат выполнения кода (порядок может меняться из-за параметра shuffle):
tensor([
[3., 4.],
[7., 8.],
]) tensor([1, 1])
tensor([
[1., 2.],
[5., 6.],
]) tensor([0, 0])
Пример
Передадим три тензора, чтобы получить кортеж из трёх элементов при обращении по индексу:
import torch
from torch.utils.data import TensorDataset
features = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
labels = torch.tensor([0, 1])
weights = torch.tensor([0.5, 0.8])
dataset = TensorDataset(features, labels, weights)
sample = dataset[1]
print(sample)
Результат выполнения кода:
(tensor([3., 4.]), tensor(1), tensor(0.8000))
Смотрите также
-
класс
IterableDataset,
который реализует набор данных для потокового чтения -
класс
ConcatDataset,
который объединяет несколько наборов данных в один -
класс
Subset,
который создаёт подмножество набора данных по индексам -
функцию
random_split,
которая разбивает набор данных на части случайным образом