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

Класс DistributedSampler

Класс DistributedSampler из модуля torch.utils.data.distributed предназначен для загрузки данных в распределенных системах обучения. Он используется совместно с DataLoader и распределяет данные между процессами таким образом, чтобы каждый процесс обрабатывал свою уникальную часть датасета.

Основная задача класса - разбить датасет на непересекающиеся части и передать каждой части соответствующему процессу. Также он обеспечивает воспроизводимость результатов за счет использования параметра seed.

Синтаксис

torch.utils.data.distributed.DistributedSampler( dataset, num_replicas=None, rank=None, shuffle=True, seed=0, drop_last=False )

Параметры

dataset (Dataset) - исходный датасет для выборки. num_replicas (int, optional) - общее количество процессов (по умолчанию берется из world_size). rank (int, optional) - номер текущего процесса (по умолчанию берется из local_rank). shuffle (bool) - перемешивать ли данные перед раздачей (по умолчанию True). seed (int) - зерно для генератора случайных чисел при перемешивании (по умолчанию 0). drop_last (bool) - отбрасывать ли последний батч, если размер датасета не делится на число процессов (по умолчанию False).

Пример использования в распределенном обучении

Рассмотрим базовый пример использования DistributedSampler в распределенном обучении. Сначала инициализируем распределенную среду:

import torch import torch.distributed as dist from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler # Инициализация распределенной среды dist.init_process_group(backend='nccl') rank = dist.get_rank() world_size = dist.get_world_size() # Создаем простой датасет class SimpleDataset(Dataset): def __init__(self): self.data = list(range(100)) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset = SimpleDataset() # Создаем DistributedSampler sampler = DistributedSampler( dataset, num_replicas=world_size, rank=rank, shuffle=True, seed=42 ) # Создаем DataLoader dataloader = DataLoader( dataset, batch_size=10, sampler=sampler ) # Выводим данные из текущего процесса for batch in dataloader: print(f"Rank {rank}, batch: {batch}") break

Пример с установкой эпохи

Для корректного перемешивания данных в каждой эпохе необходимо вызывать метод set_epoch перед каждой эпохой обучения:

import torch import torch.distributed as dist from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler dist.init_process_group(backend='nccl') rank = dist.get_rank() class SimpleDataset(Dataset): def __init__(self): self.data = list(range(50)) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset = SimpleDataset() sampler = DistributedSampler( dataset, shuffle=True, seed=123 ) dataloader = DataLoader( dataset, batch_size=5, sampler=sampler ) # Обучение в течение 2 эпох for epoch in range(2): sampler.set_epoch(epoch) # Важно для перемешивания for batch in dataloader: print(f"Epoch {epoch}, Rank {rank}, batch: {batch}")

Пример с drop_last

Параметр drop_last позволяет отбросить последний батч, если размер датасета не делится на количество процессов:

import torch from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler class SimpleDataset(Dataset): def __init__(self): self.data = list(range(23)) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset = SimpleDataset() # Создаем два процесса (для примера устанавливаем вручную) sampler_rank0 = DistributedSampler( dataset, num_replicas=2, rank=0, drop_last=True ) sampler_rank1 = DistributedSampler( dataset, num_replicas=2, rank=1, drop_last=True ) dataloader0 = DataLoader(dataset, batch_size=4, sampler=sampler_rank0) dataloader1 = DataLoader(dataset, batch_size=4, sampler=sampler_rank1) print("Rank 0 samples:") for batch in dataloader0: print(batch) print("\nRank 1 samples:") for batch in dataloader1: print(batch)

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

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