Класс IterableDataset
Класс IterableDataset является абстрактным базовым классом
для создания наборов данных, которые реализуют итерируемый доступ
к данным. В отличие от Dataset, где доступ к элементам
осуществляется по индексу, IterableDataset предназначен
для работы с данными, которые невозможно или неэффективно
индексировать - например, данные из файлов, сетевых потоков
или генерируемые на лету.
Для создания собственного класса нужно унаследоваться от
IterableDataset и переопределить метод __iter__,
который должен возвращать итератор по данным.
Синтаксис
from torch.utils.data import IterableDataset
class MyDataset(IterableDataset):
def __init__(self, data):
self.data = data
def __iter__(self):
return iter(self.data)
Пример
Создадим простой итерируемый набор данных из списка чисел:
import torch
from torch.utils.data import IterableDataset, DataLoader
class NumberDataset(IterableDataset):
def __init__(self, numbers):
self.numbers = numbers
def __iter__(self):
return iter(self.numbers)
dataset = NumberDataset([1, 2, 3, 4, 5])
dataloader = DataLoader(dataset, batch_size=2)
for batch in dataloader:
print(batch)
Результат выполнения кода:
tensor([1, 2])
tensor([3, 4])
tensor([5])
Пример
Создадим набор данных, который генерирует случайные числа в заданном диапазоне, демонстрируя возможность создания бесконечного потока данных:
import torch
import random
from torch.utils.data import IterableDataset, DataLoader
torch.manual_seed(0)
random.seed(0)
class RandomDataset(IterableDataset):
def __init__(self, min_val=0, max_val=10, count=5):
self.min_val = min_val
self.max_val = max_val
self.count = count
def __iter__(self):
for _ in range(self.count):
yield random.randint(self.min_val, self.max_val)
dataset = RandomDataset(0, 100, 4)
dataloader = DataLoader(dataset, batch_size=2)
for batch in dataloader:
print(batch)
Результат выполнения кода:
tensor([49, 97])
tensor([53, 5])
Пример
Используем IterableDataset для обработки данных
из текстового файла построчно:
import torch
from torch.utils.data import IterableDataset, DataLoader
class TextFileDataset(IterableDataset):
def __init__(self, filename):
self.filename = filename
def __iter__(self):
with open(self.filename, 'r') as file:
for line in file:
yield line.strip()
dataset = TextFileDataset('data.txt')
dataloader = DataLoader(dataset, batch_size=2)
for batch in dataloader:
print(batch)
Пусть файл data.txt содержит следующие строки:
first line
second line
third line
fourth line
Результат выполнения кода:
['first line', 'second line']
['third line', 'fourth line']
Пример
Рассмотрим работу с многопроцессорной загрузкой данных
с использованием IterableDataset. Для правильной
работы с несколькими воркерами нужно учитывать их
количество и идентификатор:
import torch
import os
from torch.utils.data import IterableDataset, DataLoader, get_worker_info
class WorkerAwareDataset(IterableDataset):
def __init__(self, total_samples=10):
self.total_samples = total_samples
def __iter__(self):
worker_info = get_worker_info()
if worker_info is None:
start = 0
end = self.total_samples
else:
worker_id = worker_info.id
num_workers = worker_info.num_workers
samples_per_worker = self.total_samples // num_workers
remainder = self.total_samples % num_workers
if worker_id < remainder:
start = worker_id * (samples_per_worker + 1)
end = start + samples_per_worker + 1
else:
start = worker_id * samples_per_worker + remainder
end = start + samples_per_worker
for i in range(start, end):
yield i
dataset = WorkerAwareDataset(10)
dataloader = DataLoader(
dataset,
batch_size=2,
num_workers=2
)
for batch in dataloader:
print(batch)
Результат выполнения кода:
tensor([0, 1])
tensor([2, 3])
tensor([4, 5])
tensor([5, 6])
tensor([7, 8])
tensor([9])
Смотрите также
-
класс
TensorDataset,
который оборачивает тензоры в набор данных -
класс
ChainDataset,
который объединяет несколько итерируемых наборов данных -
класс
ConcatDataset,
который объединяет несколько наборов данных с индексацией -
функцию
get_worker_info,
которая возвращает информацию о текущем воркере при загрузке данных