Skip to content

cyclic dataloader with data sharding repeats and skips samples when resumed at a different data-parallel size #7645

Description

@SeverinVisionary

Describe the bug

With --dataloader-type cyclic and data sharding on (the default), resuming from a checkpoint at a different data-parallel size silently repeats some samples and skips others. There is no error or warning.

The checkpoint stores only consumed_train_samples, and build_pretraining_data_loader rebuilds each rank's MegatronPretrainingRandomSampler from it. With data_sharding=True, each rank walks a permutation of its own contiguous bucket. bucket_size, bucket_offset and start_idx all depend on data_parallel_size. After a DP change, the resumed ranks walk different buckets from a different offset, so the samples they skip are not the samples the run already consumed in that epoch.

The distributed-checkpointing guide (docs/api-guide/core/dist_checkpointing.md) says a checkpoint saved under one parallel configuration, including data parallelism, can be loaded under a different one. So the model and optimizer state reshard correctly, but the data order does not.

Steps/Code to reproduce bug

This runs on CPU and needs only the sampler class (checked against main @ ec806b2).

from collections import Counter

import torch
from megatron.training.datasets.data_samplers import MegatronPretrainingRandomSampler

N, MBS, GBS = 256, 2, 16  # 16 global steps per epoch


def epoch0_draws(dp_before, dp_after, resume_step, sharding=True):
    """IDs drawn in epoch 0 when the job is resumed at `resume_step` with another DP size.
    Samplers are rebuilt from consumed_samples, as build_pretraining_data_loader does."""
    drawn = []
    for dp, lo, hi in ((dp_before, 0, resume_step), (dp_after, resume_step, N // GBS)):
        ga = GBS // (MBS * dp)
        ranks = [iter(MegatronPretrainingRandomSampler(
            torch.arange(N), total_samples=N, consumed_samples=lo * GBS, micro_batch_size=MBS,
            data_parallel_rank=r, data_parallel_size=dp, data_sharding=sharding)) for r in range(dp)]
        for _ in range(lo, hi):
            for it in ranks:
                for _ in range(ga):
                    drawn += next(it)
    c = Counter(drawn)
    return sum(v - 1 for v in c.values() if v > 1), N - len(c)


for sharding in (True, False):
    for after in (4, 2, 8):
        dup, miss = epoch0_draws(4, after, resume_step=5, sharding=sharding)
        print(f"data_sharding={sharding}, DP 4 -> {after} at step 5: {dup} duplicated, {miss} never drawn in epoch 0")

Output:

data_sharding=True, DP 4 -> 4 at step 5: 0 duplicated, 0 never drawn in epoch 0
data_sharding=True, DP 4 -> 2 at step 5: 58 duplicated, 58 never drawn in epoch 0
data_sharding=True, DP 4 -> 8 at step 5: 56 duplicated, 56 never drawn in epoch 0
data_sharding=False, DP 4 -> 4 at step 5: 0 duplicated, 0 never drawn in epoch 0
data_sharding=False, DP 4 -> 2 at step 5: 0 duplicated, 0 never drawn in epoch 0
data_sharding=False, DP 4 -> 8 at step 5: 0 duplicated, 0 never drawn in epoch 0

A sweep over every resume step and every DP change among {1, 2, 4, 8} fails at nearly every resume point, for both divisible and non-divisible dataset sizes.

Expected behavior

After a resume at any supported parallel configuration, every sample of the current epoch is drawn exactly once. Alternatively, the resume should fail loudly if the data order cannot be preserved.

Additional context

  • --dataloader-type single, the default, is not affected: it held under every DP change in the same sweep.
  • Fix cyclic sampler global microbatch boundaries #7299 (open) changes only the non-sharded path and does not affect this.
  • Possible directions, for the maintainers to choose from:
    1. Record data_parallel_size and data_sharding next to consumed_train_samples in the checkpoint. Raise a clear error, or at least a warning, when a cyclic + sharding run resumes with a different DP size.
    2. Make the sharded layout independent of the DP size. For example, bucket by a fixed shard count chosen at the start of training and stored in the checkpoint, and map ranks to buckets at load time.
    3. At minimum, document the limitation next to --dataloader-type cyclic / --no-data-sharding. Today --no-data-sharding sits in the vision argument group with the help text "Disable data sharding."

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions