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

Метод __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__,
    который определяет доступ к элементам по индексу
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить