diff --git a/megatron/training/datasets/data_samplers.py b/megatron/training/datasets/data_samplers.py index 7a31d846a49..296acc97941 100644 --- a/megatron/training/datasets/data_samplers.py +++ b/megatron/training/datasets/data_samplers.py @@ -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( 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( @@ -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. diff --git a/megatron/training/training.py b/megatron/training/training.py index 7c66812c67a..d14b7769574 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -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(): @@ -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: diff --git a/tests/unit_tests/test_training.py b/tests/unit_tests/test_training.py index a893734bd89..838b963778c 100644 --- a/tests/unit_tests/test_training.py +++ b/tests/unit_tests/test_training.py @@ -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() @@ -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: