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