Skip to content
Merged
51 changes: 46 additions & 5 deletions megatron/training/datasets/data_samplers.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,13 +41,10 @@ def build_pretraining_data_loader(dataset, consumed_samples):
global_batch_size = getattr(args, 'eval_global_batch_size', args.global_batch_size) if is_eval else args.global_batch_size

if split == Split.valid and args.full_validation:
batch_sampler = MegatronPretrainingSampler(
batch_sampler = MegatronFullValidationSampler(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we only use MegatronFullValidationSampler when required? (When using small datasets & full_validation)

Otherwise keep MegatronPretrainingSampler

Note we can access len(dataset) & DP size

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

No, we need it whenever we're doing full_validation. The wording of the comment is a little confusing though. Should I change it?

total_samples=len(dataset),
consumed_samples=0,
micro_batch_size=micro_batch_size,
data_parallel_rank=mpu.get_data_parallel_rank(),
data_parallel_size=mpu.get_data_parallel_world_size(),
)
data_parallel_size=mpu.get_data_parallel_world_size())
elif args.dataloader_type == 'single':
if args.hybrid_context_parallel:
batch_sampler = HybridCPMegatronPretrainingSampler(
Expand Down Expand Up @@ -228,6 +225,50 @@ def __iter__(self):
global_batch_idx.extend(batch[start_idx[i]:end_idx[i]])
yield global_batch_idx


class MegatronFullValidationSampler:
"""Sampler for full validation that handles small datasets gracefully.

This sampler is designed for validation datasets that may be smaller than
data_parallel_size * micro_batch_size. It uses micro_batch_size=1 to minimize
the samples needed per batch and properly handles partial batches where some
ranks may not have data.
"""

def __init__(self, total_samples, data_parallel_rank, data_parallel_size):
self.total_samples = total_samples
self.data_parallel_rank = data_parallel_rank
self.data_parallel_size = data_parallel_size
self.micro_batch_size = 1 # Always use 1 for small dataset support

# Sanity checks
assert self.total_samples > 0, f'no sample to consume: {self.total_samples}'
assert data_parallel_size > 0
assert self.data_parallel_rank < data_parallel_size, \
f'data_parallel_rank should be smaller than data size: {self.data_parallel_rank}, {data_parallel_size}'

def __len__(self):
"""Returns the number of batches this rank will yield."""
# Each batch takes data_parallel_size samples (1 per rank)
# This rank gets samples at indices: data_parallel_rank, data_parallel_rank + data_parallel_size, ...
num_batches = 0
for batch_idx in range(0, self.total_samples, self.data_parallel_size):
# Check if this rank has data in this batch
sample_idx = batch_idx + self.data_parallel_rank
if sample_idx < self.total_samples:
num_batches += 1
return num_batches

def __iter__(self):
"""Yield batches for this data parallel rank."""
for batch_idx in range(0, self.total_samples, self.data_parallel_size):
# Check if this rank has data in this batch
sample_idx = batch_idx + self.data_parallel_rank
if sample_idx < self.total_samples:
# Yield a batch with a single sample index for this rank
yield [sample_idx]


class RandomSeedDataset(Dataset):
"""
A dataset wrapper that resets the random seed before each sample.
Expand Down
46 changes: 33 additions & 13 deletions megatron/training/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -3526,10 +3526,20 @@ def evaluate_and_print_results(
print_rank_last('-' * length)


def cyclic_iter(iter):
def cyclic_iter(iterable):
while True:
for x in iter:
iterator = iter(iterable)
count = 0
for x in iterator:
count += 1
yield x
if count == 0:
# No data was yielded, the iterable is empty
raise RuntimeError(
"cyclic_iter: iterable produced no data. "
"This may indicate the validation dataloader is empty or eval_iters is incorrectly set. "
"Check that your validation dataset has data and that the dataloader is properly configured."
)


def get_train_valid_test_num_samples():
Expand Down Expand Up @@ -3704,32 +3714,42 @@ def _get_iterator(dataloader_type, dataloader):
train_data_iterator = None

if valid_dataloaders is not None:
# when using full validation, we need to override eval iters with the correct
# number of iterations on tp rank 0 so that it can be distributed to the other
# ranks later
# when using full validation, we need to override eval iters with the
# MAX length across DP ranks so all ranks run the same number of steps
if args.full_validation:
if args.multiple_validation_sets:
if valid_dataloaders[0] is None:
args.eval_iters = [None]*len(valid_dataloaders)
args.eval_iters = [None] * len(valid_dataloaders)
else:
args.eval_iters = [len(dl) for dl in valid_dataloaders]
local_eval_iters = [len(dl) for dl in valid_dataloaders]
eval_iters_tensor = torch.tensor(local_eval_iters, dtype=torch.long, device='cuda')
torch.distributed.all_reduce(
eval_iters_tensor,
op=torch.distributed.ReduceOp.MAX,
group=mpu.get_data_parallel_group(with_context_parallel=True),
)
args.eval_iters = eval_iters_tensor.tolist()
else:
args.eval_iters = len(valid_dataloaders[0])
local_eval_iters = len(valid_dataloaders[0])
eval_iters_tensor = torch.tensor([local_eval_iters], dtype=torch.long, device='cuda')
torch.distributed.all_reduce(
eval_iters_tensor,
op=torch.distributed.ReduceOp.MAX,
group=mpu.get_data_parallel_group(with_context_parallel=True),
)
args.eval_iters = eval_iters_tensor.item()

if args.multiple_validation_sets:
if valid_dataloaders[0] is None:
valid_data_iterators = [None] * len(valid_dataloaders)
else:
valid_dl_type = "cyclic" if args.full_validation else dl_type
print(
f"[VALID DATA LOADER LENGTHS] "
", ".join(f"{idx}: {len(dl)}" for idx, dl in enumerate(valid_dataloaders))
)
valid_data_iterators = [
_get_iterator(valid_dl_type, dl) for dl in valid_dataloaders
]
elif valid_dataloaders[0] is not None:
valid_data_iterators = _get_iterator(dl_type, valid_dataloaders[0])
valid_dl_type = "cyclic" if args.full_validation else dl_type
valid_data_iterators = _get_iterator(valid_dl_type, valid_dataloaders[0])
else:
valid_data_iterators = None
else:
Expand Down
36 changes: 36 additions & 0 deletions tests/unit_tests/test_training.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,24 @@ def mock_train_valid_test_datasets_provider(train_val_test_num_samples):
return iter([1]), iter([2]), iter([3])


class _LenDataloader:
"""Fake dataloader with __len__ (required by the full_validation path)
and __iter__ (consumed via cyclic_iter)."""

def __init__(self, data):
self._data = list(data)

def __len__(self):
return len(self._data)

def __iter__(self):
return iter(self._data)


def mock_multi_valid_full_datasets_provider(train_val_test_num_samples):
return (iter([1]), [_LenDataloader([2, 2]), _LenDataloader([20, 20, 20])], iter([3]))


def create_test_args():
# Set dummy values for the args.
args = SimpleNamespace()
Expand Down Expand Up @@ -55,6 +73,24 @@ def test_build_train_valid_test_data_iterators(self):
test_data = next(test_iter)
assert (train_data, valid_data, test_data) == (1, 2, 3)

def test_build_train_valid_test_data_iterators_multi_full_validation(self):
"""multiple_validation_sets + full_validation builds a list of iterators
(one per validation set) and sets args.eval_iters to the per-loader
lengths MAX-reduced across DP ranks."""
args = create_test_args()
args.multiple_validation_sets = True
args.full_validation = True
set_args(args)
_, valid_iters, _ = build_train_valid_test_data_iterators(
mock_multi_valid_full_datasets_provider
)
assert isinstance(valid_iters, list)
assert len(valid_iters) == 2
assert next(valid_iters[0]) == 2
assert next(valid_iters[1]) == 20
# data_parallel_size=1, so MAX across DP ranks equals the local lengths
assert args.eval_iters == [2, 3]

def test_closed_formula_vocab_size_with_padding(self):
def old_round_impl(after, multiple):
while (after % multiple) != 0:
Expand Down
Loading