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

Метод __add__ класса Dataset

Метод __add__ класса Dataset позволяет объединять два датасета в один. В результате работы метода возвращается новый объект типа ConcatDataset, который содержит элементы обоих исходных датасетов. Метод вызывается при использовании оператора сложения + между двумя датасетами.

Синтаксис

dataset = dataset1 + dataset2

Пример

Давайте создадим два простых датасета и объединим их с помощью метода __add__:

import torch from torch.utils.data import Dataset class SimpleDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return torch.tensor(self.data[idx]) dataset1 = SimpleDataset([1, 2, 3, 4, 5]) dataset2 = SimpleDataset([6, 7, 8, 9, 10]) combined = dataset1 + dataset2 print(len(combined))

Результат выполнения кода:

10

Пример

Давайте проверим, что элементы объединённого датасета соответствуют элементам исходных датасетов:

import torch from torch.utils.data import Dataset class SimpleDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return torch.tensor(self.data[idx]) dataset1 = SimpleDataset([1, 2, 3]) dataset2 = SimpleDataset([4, 5, 6]) combined = dataset1 + dataset2 for i in range(len(combined)): print(combined[i])

Результат выполнения кода:

tensor(1) tensor(2) tensor(3) tensor(4) tensor(5) tensor(6)

Пример

Давайте объединим три датасета с помощью цепочки операций сложения:

import torch from torch.utils.data import Dataset class SimpleDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return torch.tensor(self.data[idx]) dataset1 = SimpleDataset([1, 2]) dataset2 = SimpleDataset([3, 4]) dataset3 = SimpleDataset([5, 6]) combined = dataset1 + dataset2 + dataset3 print(len(combined)) print(combined[0]) print(combined[3]) print(combined[5])

Результат выполнения кода:

6 tensor(1) tensor(4) tensor(6)

Смотрите также

  • класс Dataset,
    который является базовым классом для создания пользовательских датасетов
  • метод __getitem__,
    который возвращает элемент датасета по индексу
  • метод __len__,
    который возвращает размер датасета
  • класс ConcatDataset,
    который используется для объединения нескольких датасетов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить