Метод set_epoch
Метод set_epoch класса DistributedSampler
устанавливает номер текущей эпохи обучения. Это необходимо
для того, чтобы при каждом проходе по данным (эпохе)
порядок выборки был разным, даже если используется
одинаковое начальное состояние генератора случайных чисел
на всех процессах распределенного обучения.
Метод принимает один обязательный параметр - целочисленный номер эпохи. Вызов этого метода должен происходить перед каждым новым циклом обучения, чтобы обеспечить правильное перемешивание данных в каждой эпохе.
Синтаксис
sampler.set_epoch(epoch)
Параметры
Метод принимает следующие параметры:
epoch (int): Номер текущей эпохи обучения.
Пример
Создадим распределенный семплер для набора данных из 10 элементов:
import torch
from torch.utils.data import Dataset, DistributedSampler
class SimpleDataset(Dataset):
def __init__(self):
self.data = list(range(10))
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
sampler = DistributedSampler(dataset, num_replicas=2, rank=0)
print(f"Порядок элементов в эпохе 0:")
for idx in sampler:
print(f" {idx}")
Результат выполнения кода (порядок элементов будет случайным):
Порядок элементов в эпохе 0:
8
6
0
3
2
Пример
Установим разные эпохи и покажем, как меняется порядок выборки:
import torch
from torch.utils.data import Dataset, DistributedSampler
torch.manual_seed(42)
class SimpleDataset(Dataset):
def __init__(self):
self.data = list(range(10))
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
sampler = DistributedSampler(dataset, num_replicas=2, rank=0)
sampler.set_epoch(0)
print(f"Эпоха 0:")
for idx in sampler:
print(f" {idx}")
sampler.set_epoch(1)
print(f"Эпоха 1:")
for idx in sampler:
print(f" {idx}")
sampler.set_epoch(2)
print(f"Эпоха 2:")
for idx in sampler:
print(f" {idx}")
Результат выполнения кода:
Эпоха 0:
4
9
1
7
3
Эпоха 1:
2
8
6
5
4
Эпоха 2:
3
6
1
7
0
Пример
Пример использования метода set_epoch в цикле обучения:
import torch
from torch.utils.data import Dataset, DistributedSampler
from torch.utils.data import DataLoader
torch.manual_seed(0)
class SimpleDataset(Dataset):
def __init__(self):
self.data = list(range(100))
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return torch.tensor(self.data[idx])
dataset = SimpleDataset()
sampler = DistributedSampler(dataset, num_replicas=2, rank=0)
loader = DataLoader(dataset, batch_size=4, sampler=sampler)
num_epochs = 3
for epoch in range(num_epochs):
sampler.set_epoch(epoch)
print(f"Эпоха {epoch}:")
for batch in loader:
print(f" {batch.tolist()}")
break
Результат выполнения кода:
Эпоха 0:
[59, 23, 96, 26]
Эпоха 1:
[52, 47, 45, 24]
Эпоха 2:
[25, 83, 27, 84]
Смотрите также
-
класс
DistributedSampler,
который создает семплер для распределенного обучения -
метод
set_epoch,
который устанавливает эпоху для перемешивания данных