Сэмплер по рангам в PyTorch
Чтобы процессы не брали
одни и те же строки,
используют DistributedSampler.
Он знает число копий
обучения и ранг текущего
процесса, даже если группа
задана только числами.
Зададим два процесса и ранг ноль на наборе из пяти элементов и выведем, сколько индексов достанется этому рангу:
import torch
from torch.utils.data import TensorDataset, DistributedSampler
rows = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
dataset = TensorDataset(rows)
sampler = DistributedSampler(
dataset, num_replicas=2, rank=0
)
print(len(sampler)) # выведет 3
Перед каждой эпохой сэмплер перемешивает индексы по своему правилу. В загрузчик его передают вместо обычной перестановки, когда обучение идёт на нескольких процессах.
Для набора из 6
строк, двух копий и
ранга 1 постройте
сэмплер по рангам и
выведите длину выборки
для этого ранга.
На наборе из 4
элементов с двумя
копиями сравните длины
выборки для рангов
0 и 1.
Выведите обе длины
через пробел.
Для пяти чисел и трёх
копий с рангом 2
создайте сэмплер и
выведите длину его
выборки.