Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Distributed sampler tails and global loss weighting

Last updated: 6 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

Data-parallel workers must agree on which examples count and on the denominator used for one optimizer update; an even number of steps alone does not guarantee that.

Count sampled identities before scaling

A distributed sampler partitions dataset indices across ranks. When the dataset length is not divisible by worker count, one policy repeats indices to equalize shard lengths; another drops a tail. With 47 unique receipts and three workers, padding to 48 slots repeats one receipt. That extra slot changes the training population unless it is explicitly masked or accepted as the sampling policy. Audit unique identities, duplicate identities and omitted identities per epoch. Sampling choices can outweigh a small optimizer change.

Refresh order every epoch

A distributed sampler seeded once can generate the same shuffled order each epoch unless its epoch value is advanced before the loader iterator is created. The epoch value must be the same on all ranks. If one worker skips an epoch or has a different dataset manifest, ranks may see different step counts and collective communication can hang. Log world size, rank, manifest hash, sampler seed, epoch and accepted-example count. A successful first epoch does not prove that later epochs reshuffle.

Match the global denominator

A conventional data-parallel reducer averages gradients across workers. If all ranks have the same local batch size and each computes a local mean, that average is the global example mean. If rank zero has three examples while rank one has one, averaging their local means gives each rank equal weight, not each example. For a sample-averaged loss, each rank should contribute its local loss sum multiplied by world size and divided by the global accepted-example count before the gradient average. Agree on any sampler duplicates and ignored labels first. Effective-batch accounting extends across ranks.

Handle uneven work deliberately

A worker that runs out of data earlier can leave peers waiting at an all-reduce. Either construct equal work with a documented duplicate or drop policy, or use a supported uneven-input join mechanism with a stated gradient-weight convention. Do not let a broad exception handler swallow an exhausted-loader event while other ranks continue. For validation, avoid padding duplicates in the metric denominator; gather predictions with stable example identities and deduplicate before calculating reported rates.

Prove parity on a tiny batch

Copy one model to two simulated workers. Backpropagate four examples as one batch, then split them three and one, apply the distributed scaling rule and average worker gradients. Compare every parameter gradient before stepping. Repeat with a zero-valid-target rank, a short final step and a sampler epoch change. The code example below audits sampler slots; the paired project performs the gradient comparison. A low training loss alone cannot reveal a wrong denominator.

Implementation

python
from collections import Counter
from torch.utils.data import DistributedSampler, TensorDataset
import torch

receipt_features = torch.arange(47, dtype=torch.float32).unsqueeze(1)
receipt_dataset = TensorDataset(receipt_features)
worker_count = 3
samplers = [DistributedSampler(receipt_dataset, num_replicas=worker_count,
                               rank=worker_rank, shuffle=True, seed=71,
                               drop_last=False)
            for worker_rank in range(worker_count)]

for sampler in samplers:
    sampler.set_epoch(2)
second_epoch_slots = [receipt_id for sampler in samplers
                      for receipt_id in sampler]
slot_counts = Counter(second_epoch_slots)
assert len(second_epoch_slots) == 48
assert len(slot_counts) == 47
assert sum(count - 1 for count in slot_counts.values()) == 1

for sampler in samplers:
    sampler.set_epoch(3)
third_epoch_slots = [receipt_id for sampler in samplers
                     for receipt_id in sampler]
assert second_epoch_slots != third_epoch_slots

Performance and operating cost

Data parallelism replicates model parameters and optimizer state on each worker. For P parameters, gradient synchronization communicates on the order of P values per update, although bucket overlap and network topology affect wall time. Sampler bookkeeping is cheap relative to training, but repeated tail examples alter the effective data weights. Scaling by the wrong denominator can be fast and still train a different objective.

Common Mistakes

  • Do not assume a non-divisible dataset is partitioned without duplicates.
  • Do not advance the sampler seed independently on each worker.
  • Do not average unequal local means when the intended unit is one accepted example.

Read next

Continue the workflow: Contrastive temperature and negative-mask accounting.

Continue the workflow: Expert load balance, overflow slices and dispatch consistency.

ai-data
deep-learning
Storage details