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

Класс ChainDataset

Класс ChainDataset из модуля torch.utils.data применяется для последовательного объединения нескольких объектов IterableDataset в один сквозной датасет. При итерации по такому датасету данные сначала перебираются из первого датасета, затем из второго и так далее. Это полезно, когда нужно обработать данные, которые естественным образом разбиты на несколько независимых частей, но для модели они должны подаваться как единый поток.

Первым и единственным параметром конструктор принимает итерируемый объект, элементами которого являются экземпляры IterableDataset.

Синтаксис

torch.utils.data.ChainDataset(datasets)

Пример

Давайте создадим два простых итерируемых датасета и объединим их в цепочку:

import torch from torch.utils.data import IterableDataset, ChainDataset class MyIterableDataset(IterableDataset): def __init__(self, data): self.data = data def __iter__(self): return iter(self.data) dataset1 = MyIterableDataset([1, 2, 3]) dataset2 = MyIterableDataset([4, 5, 6]) chain = ChainDataset([dataset1, dataset2]) for item in chain: print(item)

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

1 2 3 4 5 6

Пример

Объединим три датасета, содержащих числа разных диапазонов, и получим их сумму:

import torch from torch.utils.data import IterableDataset, ChainDataset class RangeDataset(IterableDataset): def __init__(self, start, end): self.start = start self.end = end def __iter__(self): return iter(range(self.start, self.end)) dataset1 = RangeDataset(0, 3) dataset2 = RangeDataset(10, 13) dataset3 = RangeDataset(20, 23) chain = ChainDataset([dataset1, dataset2, dataset3]) total = 0 for item in chain: total += item print(total)

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

66

Пример

Рассмотрим работу ChainDataset с датасетами, которые генерируют бесконечные последовательности. В этом случае важно ограничить количество шагов итерации, чтобы не создать бесконечный цикл:

import torch from torch.utils.data import IterableDataset, ChainDataset class InfiniteDataset(IterableDataset): def __init__(self, multiplier): self.multiplier = multiplier def __iter__(self): i = 0 while True: yield i * self.multiplier i += 1 dataset1 = InfiniteDataset(1) dataset2 = InfiniteDataset(10) chain = ChainDataset([dataset1, dataset2]) res = [] for idx, item in enumerate(chain): if idx == 5: break res.append(item) print(res)

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

[0, 1, 2, 3, 4, 0]

В этом примере сначала берутся пять элементов из первого бесконечного датасета, а затем один элемент из второго. Обратите внимание, что порядок следования датасетов в списке определяет порядок их перебора.

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

  • класс IterableDataset,
    который позволяет создавать пользовательские итерируемые датасеты
  • класс ConcatDataset,
    который объединяет датасеты, поддерживающие индексацию
  • класс TensorDataset,
    который позволяет работать с тензорами как с датасетами
  • функцию default_collate,
    которая используется для сборки батчей при загрузке данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить