Свой класс набора данных в PyTorch
Если данные хранятся в полях объекта,
набор описывают классом от
Dataset из
torch.utils.data. Нужны
метод __len__, который
сообщает число элементов, и метод
__getitem__, который по
индексу возвращает пару тензоров.
Опишем класс с двумя строками признаков и двумя метками, создадим объект и проверим длину и один элемент:
import torch
from torch.utils.data import Dataset
class PairTable(Dataset):
def __init__(self):
self.features = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
self.targets = torch.tensor([0, 1])
def __len__(self):
return len(self.features)
def __getitem__(self, index):
return self.features[index], self.targets[index]
ds = PairTable()
print(len(ds)) # выведет 2
print(ds[0]) # выведет (tensor([1., 2.]), tensor(0))
Внутри класса тензоры могут собираться из файлов или считаться на лету; снаружи остаётся тот же контракт: длина и выдача пары по номеру.
Опишите класс с тремя строками
[[1.0], [2.0], [3.0]]
и метками [10, 20, 30].
Создайте объект и выведите,
сколько в нём элементов.
Опишите класс с двумя строками
по три числа и метками 0
и 1. Возьмите элемент
с индексом 1 и выведите
его метку.
Опишите класс, где признаки -
таблица [[0.0, 0.0], [1.0, 1.0]],
а метки - 0 и 1.
Выведите признаки первого
элемента.