Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
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 для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить