РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
⊗pytoPmDtDs 89 of 95 menu
◀ ▶

Свой класс набора данных в 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. Выведите признаки первого элемента.

← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить