Класс 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,
которая используется для сборки батчей при загрузке данных