Класс 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,
которая инициализирует распределенную среду