Метод __getitem__
Метод __getitem__ является ключевым методом
класса Dataset в PyTorch. Он определяет,
как именно извлекается элемент из датасета по
заданному индексу. При создании собственного
класса датасета этот метод обязательно должен
быть переопределён. В качестве первого параметра
метод принимает индекс idx, а возвращает
кортеж из тензора с данными и тензора с меткой
(или другой информации о примере).
Синтаксис
class MyDataset(Dataset):
def __getitem__(self, idx):
# загрузка данных по индексу
# возврат данных и метки
Пример
Давайте создадим простой датасет с числами
от 0 до 4 и их квадратами:
import torch
from torch.utils.data import Dataset
class SimpleDataset(Dataset):
def __init__(self):
self.data = [0, 1, 2, 3, 4]
self.labels = [0, 1, 4, 9, 16]
def __getitem__(self, idx):
x = torch.tensor(self.data[idx])
y = torch.tensor(self.labels[idx])
return x, y
def __len__(self):
return len(self.data)
dataset = SimpleDataset()
item = dataset[2]
print(item)
Результат выполнения кода:
(tensor(2), tensor(4))
Пример
Теперь создадим датасет, который загружает данные из списка словарей и возвращает их в виде тензоров с плавающей точкой:
import torch
from torch.utils.data import Dataset
class DictDataset(Dataset):
def __init__(self, records):
self.records = records
def __getitem__(self, idx):
record = self.records[idx]
x = torch.tensor(record['features'], dtype=torch.float)
y = torch.tensor(record['target'], dtype=torch.long)
return x, y
def __len__(self):
return len(self.records)
records = [
{'features': [1.0, 2.0], 'target': 0},
{'features': [3.0, 4.0], 'target': 1},
{'features': [5.0, 6.0], 'target': 0},
]
dataset = DictDataset(records)
item = dataset[1]
print(item)
Результат выполнения кода:
(tensor([3., 4.]), tensor(1))
Пример
Давайте реализуем датасет, который возвращает не только данные и метку, но и дополнительную информацию в виде словаря:
import torch
from torch.utils.data import Dataset
class ExtendedDataset(Dataset):
def __init__(self, data):
self.data = data
def __getitem__(self, idx):
item = self.data[idx]
x = torch.tensor(item['input'], dtype=torch.float)
y = torch.tensor(item['output'], dtype=torch.float)
info = {
'id': item['id'],
'name': item['name'],
}
return x, y, info
def __len__(self):
return len(self.data)
data = [
{'id': 1, 'name': 'sample1', 'input': [0.1, 0.2], 'output': [0.3]},
{'id': 2, 'name': 'sample2', 'input': [0.4, 0.5], 'output': [0.6]},
]
dataset = ExtendedDataset(data)
res = dataset[0]
print(res[0])
print(res[1])
print(res[2])
Результат выполнения кода:
tensor([0.1000, 0.2000])
tensor([0.3000])
{'id': 1, 'name': 'sample1'}
Пример
Рассмотрим случай, когда датасет работает с
изображениями, хранящимися в виде тензоров
(например, размерность 3x32x32):
import torch
from torch.utils.data import Dataset
class ImageDataset(Dataset):
def __init__(self, images, labels):
self.images = images
self.labels = labels
def __getitem__(self, idx):
img = torch.tensor(self.images[idx], dtype=torch.float)
label = torch.tensor(self.labels[idx], dtype=torch.long)
return img, label
def __len__(self):
return len(self.images)
torch.manual_seed(0)
images = torch.randn(5, 3, 32, 32).tolist()
labels = [0, 1, 0, 1, 0]
dataset = ImageDataset(images, labels)
img, label = dataset[3]
print(img.shape)
print(label)
Результат выполнения кода:
torch.Size([3, 32, 32])
tensor(1)
Смотрите также
-
класс
Dataset,
который является базовым классом для всех датасетов -
метод
__len__,
который возвращает размер датасета -
метод
__add__,
который объединяет два датасета -
метод
__getitem__,
который определяет доступ к элементам по индексу