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

Атрибут generator

Атрибут generator класса DataLoader задает генератор случайных чисел, который используется для перемешивания данных при каждом проходе эпохи. Этот атрибут позволяет контролировать случайность процесса загрузки данных и обеспечивать воспроизводимость результатов.

Синтаксис

torch.utils.data.DataLoader( dataset, batch_size=1, shuffle=False, sampler=None, batch_sampler=None, num_workers=0, collate_fn=None, pin_memory=False, drop_last=False, timeout=0, worker_init_fn=None, multiprocessing_context=None, generator=None, *, prefetch_factor=2, persistent_workers=False, pin_memory_device='' )

Атрибут можно установить через параметр generator при создании объекта DataLoader.

Пример

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

import torch from torch.utils.data import DataLoader, TensorDataset torch.manual_seed(0) data = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) dataset = TensorDataset(data) generator = torch.Generator() generator.manual_seed(42) dataloader1 = DataLoader( dataset, batch_size=2, shuffle=True, generator=generator ) dataloader2 = DataLoader( dataset, batch_size=2, shuffle=True, generator=generator ) res1 = [batch.item() for batch, in dataloader1] res2 = [batch.item() for batch, in dataloader2] print(res1) print(res2)

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

[6, 4, 10, 7, 5, 2, 9, 8, 3, 1] [6, 4, 10, 7, 5, 2, 9, 8, 3, 1]

Пример

Создадим два загрузчика с разными генераторами и сравним порядок выборки данных:

import torch from torch.utils.data import DataLoader, TensorDataset torch.manual_seed(0) data = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) dataset = TensorDataset(data) generator1 = torch.Generator() generator1.manual_seed(42) generator2 = torch.Generator() generator2.manual_seed(123) dataloader1 = DataLoader( dataset, batch_size=2, shuffle=True, generator=generator1 ) dataloader2 = DataLoader( dataset, batch_size=2, shuffle=True, generator=generator2 ) res1 = [batch.item() for batch, in dataloader1] res2 = [batch.item() for batch, in dataloader2] print(res1) print(res2)

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

[6, 4, 10, 7, 5, 2, 9, 8, 3, 1] [4, 3, 10, 8, 7, 6, 2, 5, 1, 9]

Пример

При работе с многопоточностью важно установить генератор для обеспечения воспроизводимости результатов:

import torch from torch.utils.data import DataLoader, TensorDataset torch.manual_seed(0) data = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) dataset = TensorDataset(data) generator = torch.Generator() generator.manual_seed(42) dataloader = DataLoader( dataset, batch_size=3, shuffle=True, num_workers=2, generator=generator ) res = [batch.item() for batch, in dataloader] print(res)

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

[6, 4, 10, 7, 5, 2, 9, 8, 3, 1]

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

  • класс DataLoader,
    который загружает данные для обучения моделей
  • атрибут dataset,
    который содержит исходный набор данных
  • атрибут batch_size,
    который определяет размер батча при загрузке
  • атрибут sampler,
    который определяет стратегию выборки данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить