分布式sampler如何保证每个epoch的数据不重复?
多卡数据并行时,每张卡用不同的数据子集,且整个 epoch 内所有卡合起来正好把数据集覆盖一遍、不重不漏。DistributedSampler 就是按 rank、world_size 把样本划分到各卡,并可选打乱(shuffle)后每卡只取自己的那一段。
n_per_rank = ceil(N/W),总长度可能补到 n_per_rank * W(不足用重复或 drop)。卡 rank 拿的下标为 rank, rank+W, rank+2W, ...,即按 rank 交错,这样每卡拿到不重叠的一批下标。shuffle=True,每个 epoch 开始时对「全局下标 0..N-1」做一次 shuffle(用相同的 seed,如 epoch),再按上面规则按 rank 取;这样每 epoch 每卡看到的顺序不同,但卡间仍不重叠。关键:所有进程用同一 seed(例如 sampler.set_epoch(epoch) 里用 epoch 作 seed),这样每卡上的 shuffle 结果一致,再按 rank 切分后仍不重不漏。sampler.set_epoch(epoch),让 shuffle 的 seed 随 epoch 变,否则每个 epoch 每卡拿到的是同一顺序,可能影响收敛。| 返回模块 | 返回总览 |