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

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