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

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