From cf6fe96b9480d2d7f2838d5d7336f154c50c8d6c Mon Sep 17 00:00:00 2001 From: cx Date: Mon, 10 Aug 2026 09:47:38 +0000 Subject: [PATCH] fix(dataloader): track distributed progress by global batch --- example/gpt2/main.cc | 34 +++++----- example/llama3/main.cc | 34 +++++----- infini_train/include/dataloader.h | 13 ++-- infini_train/src/dataloader.cc | 75 ++++++++++++++++----- tests/CMakeLists.txt | 3 + tests/dataloader/CMakeLists.txt | 8 +++ tests/dataloader/test_dataloader.cc | 101 ++++++++++++++++++++++++++++ 7 files changed, 208 insertions(+), 60 deletions(-) create mode 100644 tests/dataloader/CMakeLists.txt create mode 100644 tests/dataloader/test_dataloader.cc diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index fd2522e3..988971d6 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -372,15 +372,19 @@ void Train(const nn::parallel::Rank &rank) { start_step = resume_result.global_step; size_t consumed_batches = resume_result.consumed_batches; - // TODO(jym): Replace with Sampler abstraction when available. - // Skip dataloader to resume from the correct batch position. - if (consumed_batches > 0) { - size_t start = train_iter.BatchIndex(); - // Each rank processes every ddp_world_size-th batch starting from its own rank. - // num_skips calculates how many ++ iterations to reach the saved batch position. - size_t num_skips = (consumed_batches - start) / ddp_world_size; - for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } - } + // consumed_batches is the number of global dataloader batches consumed across cyclic epochs. + train_iter.SeekGlobalBatch(consumed_batches % train_loader.NumGlobalBatches()); + auto next_train_batch = [&]() { + auto batch = *train_iter; + // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below + // TODO(dcj): support dataloader.reset() later + ++train_iter; + if (train_iter == train_loader.end()) { + train_iter = train_loader.begin(); + } + ++consumed_batches; + return batch; + }; auto save_checkpoint = [&](const std::filesystem::path &save_dir, int64_t global_step) { SaveCheckpoint({ @@ -454,11 +458,7 @@ void Train(const nn::parallel::Rank &rank) { infini_train::AutocastGuard autocast_guard(device.type(), dtype); // (bs, seq_len), (bs, seq_len) - auto [x, y] = *train_iter; - // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below - // TODO(dcj): support dataloader.reset() later - ++train_iter; - consumed_batches = train_iter.BatchIndex(); + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -488,11 +488,7 @@ void Train(const nn::parallel::Rank &rank) { scheduler->Step(); } } else { - auto [x, y] = *train_iter; - // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below - // TODO(dcj): support dataloader.reset() later - ++train_iter; - consumed_batches = train_iter.BatchIndex(); + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 05fad4a8..14180d8b 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -354,15 +354,19 @@ void Train(const nn::parallel::Rank &rank) { start_step = resume_result.global_step; size_t consumed_batches = resume_result.consumed_batches; - // TODO(jym): Replace with Sampler abstraction when available. - // Skip dataloader to resume from the correct batch position. - if (consumed_batches > 0) { - size_t start = train_iter.BatchIndex(); - // Each rank processes every ddp_world_size-th batch starting from its own rank. - // num_skips calculates how many ++ iterations to reach the saved batch position. - size_t num_skips = (consumed_batches - start) / ddp_world_size; - for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } - } + // consumed_batches is the number of global dataloader batches consumed across cyclic epochs. + train_iter.SeekGlobalBatch(consumed_batches % train_loader.NumGlobalBatches()); + auto next_train_batch = [&]() { + auto batch = *train_iter; + // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below + // TODO(dcj): support dataloader.reset() later + ++train_iter; + if (train_iter == train_loader.end()) { + train_iter = train_loader.begin(); + } + ++consumed_batches; + return batch; + }; auto save_checkpoint = [&](const std::filesystem::path &save_dir, int64_t global_step) { SaveCheckpoint({ @@ -434,11 +438,7 @@ void Train(const nn::parallel::Rank &rank) { infini_train::AutocastGuard autocast_guard(device.type(), dtype); // (bs, seq_len), (bs, seq_len) - auto [x, y] = *train_iter; - // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below - // TODO(dcj): support dataloader.reset() later - ++train_iter; - consumed_batches = train_iter.BatchIndex(); + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -467,11 +467,7 @@ void Train(const nn::parallel::Rank &rank) { scheduler->Step(); } } else { - auto [x, y] = *train_iter; - // if we are trying to overfit a single batch, we reset the loader here by commenting out the line below - // TODO(dcj): support dataloader.reset() later - ++train_iter; - consumed_batches = train_iter.BatchIndex(); + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); diff --git a/infini_train/include/dataloader.h b/infini_train/include/dataloader.h index 38fa02a4..34dd04b7 100644 --- a/infini_train/include/dataloader.h +++ b/infini_train/include/dataloader.h @@ -12,7 +12,7 @@ class Tensor; namespace infini_train { class DataLoaderIterator { public: - DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t batch_idx, size_t max_batch_idx, + DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t global_batch_idx, size_t num_global_batches, size_t ddp_rank = 0, size_t ddp_world_size = 1); std::pair, std::shared_ptr> operator*() const; @@ -24,13 +24,14 @@ class DataLoaderIterator { friend bool operator!=(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs); friend bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs); - size_t BatchIndex() const; + size_t GlobalBatchIndex() const; + DataLoaderIterator &SeekGlobalBatch(size_t global_batch_idx); private: const Dataset *dataset_ = nullptr; // not owned size_t batch_size_ = 0; - size_t batch_idx_ = 0; - size_t max_batch_idx_ = 0; + size_t global_batch_idx_ = 0; + size_t num_global_batches_ = 0; size_t ddp_rank_ = 0; size_t ddp_world_size_ = 1; }; @@ -42,10 +43,12 @@ class DataLoader { virtual DataLoaderIterator begin() const; virtual DataLoaderIterator end() const; + size_t NumGlobalBatches() const; + protected: std::shared_ptr dataset_; size_t batch_size_ = 0; - size_t max_batch_idx_ = 0; + size_t num_global_batches_ = 0; }; class DistributedDataLoader : public DataLoader { diff --git a/infini_train/src/dataloader.cc b/infini_train/src/dataloader.cc index b7cc94f2..c7f64b41 100644 --- a/infini_train/src/dataloader.cc +++ b/infini_train/src/dataloader.cc @@ -3,6 +3,7 @@ #include #include #include +#include #include #include @@ -13,8 +14,14 @@ namespace infini_train { namespace { +size_t CheckedCeilDiv(size_t numerator, size_t denominator) { + CHECK_GT(denominator, 0); + return (numerator + denominator - 1) / denominator; +} + // TODO(dcj): Use official stack implementation later. std::shared_ptr Stack(const std::vector> &tensors) { + CHECK(!tensors.empty()) << "Cannot stack an empty batch. Check DataLoader iterator end handling."; const int batch_size = tensors.size(); const auto &dims = tensors[0]->Dims(); const int stacked_dim = std::accumulate(dims.begin(), dims.end(), 1, std::multiplies()); @@ -35,21 +42,31 @@ std::shared_ptr Stack(const std::vector> &tensor } } // namespace -DataLoaderIterator::DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t batch_idx, - size_t max_batch_idx, size_t ddp_rank, size_t ddp_world_size) - : dataset_(&dataset), batch_size_(batch_size), batch_idx_(batch_idx), max_batch_idx_(max_batch_idx), - ddp_rank_(ddp_rank), ddp_world_size_(ddp_world_size){}; +DataLoaderIterator::DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t global_batch_idx, + size_t num_global_batches, size_t ddp_rank, size_t ddp_world_size) + : dataset_(&dataset), batch_size_(batch_size), global_batch_idx_(global_batch_idx), + num_global_batches_(num_global_batches), ddp_rank_(ddp_rank), ddp_world_size_(ddp_world_size){}; std::pair, std::shared_ptr> DataLoaderIterator::operator*() const { /* 0, 1, ..., x, ... [0, bs-1], [bs, 2*bs-1], ..., [x*bs, (x+1)*bs-1], ... ^ - batch_idx + global_batch_idx */ std::vector> data_vec; std::vector> label_vec; - for (int idx = batch_idx_ * batch_size_; idx < (batch_idx_ + 1) * batch_size_ && idx < dataset_->Size(); ++idx) { + CHECK_LT(global_batch_idx_, num_global_batches_) + << "Cannot dereference DataLoader end iterator. global_batch_idx=" << global_batch_idx_ + << ", num_global_batches=" << num_global_batches_ << ", ddp_rank=" << ddp_rank_ + << ", ddp_world_size=" << ddp_world_size_; + const size_t start_idx = (global_batch_idx_ * ddp_world_size_ + ddp_rank_) * batch_size_; + CHECK_LT(start_idx, dataset_->Size()) + << "DataLoader batch starts past dataset end. global_batch_idx=" << global_batch_idx_ + << ", start_idx=" << start_idx << ", dataset_size=" << dataset_->Size() << ", batch_size=" << batch_size_ + << ", ddp_rank=" << ddp_rank_ << ", ddp_world_size=" << ddp_world_size_; + const size_t end_idx = std::min(start_idx + batch_size_, dataset_->Size()); + for (size_t idx = start_idx; idx < end_idx; ++idx) { auto &&[data, label] = dataset_->operator[](idx); data_vec.push_back(std::move(data)); label_vec.push_back(std::move(label)); @@ -58,7 +75,7 @@ std::pair, std::shared_ptr> DataLoaderIterator:: } DataLoaderIterator &DataLoaderIterator::operator++() { - batch_idx_ = std::min(batch_idx_ + ddp_world_size_, max_batch_idx_); + global_batch_idx_ = std::min(global_batch_idx_ + 1, num_global_batches_); return *this; } @@ -68,36 +85,60 @@ DataLoaderIterator DataLoaderIterator::operator++(int) { return tmp; } -bool operator<(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { return lhs.batch_idx_ < rhs.batch_idx_; } +bool operator<(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { + return lhs.global_batch_idx_ < rhs.global_batch_idx_; +} bool operator!=(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { - return lhs.batch_idx_ != rhs.batch_idx_; + return lhs.global_batch_idx_ != rhs.global_batch_idx_; } bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { - return lhs.batch_idx_ == rhs.batch_idx_; + return lhs.global_batch_idx_ == rhs.global_batch_idx_; } -size_t DataLoaderIterator::BatchIndex() const { return batch_idx_; } +size_t DataLoaderIterator::GlobalBatchIndex() const { return global_batch_idx_; } + +DataLoaderIterator &DataLoaderIterator::SeekGlobalBatch(size_t global_batch_idx) { + CHECK_LE(global_batch_idx, num_global_batches_) + << "Cannot seek past DataLoader end. global_batch_idx=" << global_batch_idx + << ", num_global_batches=" << num_global_batches_; + global_batch_idx_ = global_batch_idx; + return *this; +} DataLoader::DataLoader(const std::shared_ptr &dataset, size_t batch_size) - : dataset_(dataset), batch_size_(batch_size), max_batch_idx_((dataset_->Size() + batch_size_ - 1) / batch_size_) {} + : dataset_(dataset), batch_size_(batch_size), num_global_batches_(CheckedCeilDiv(dataset_->Size(), batch_size_)) {} -DataLoaderIterator DataLoader::begin() const { return DataLoaderIterator(*dataset_, batch_size_, 0, max_batch_idx_); } +DataLoaderIterator DataLoader::begin() const { + return DataLoaderIterator(*dataset_, batch_size_, 0, num_global_batches_); +} DataLoaderIterator DataLoader::end() const { - return DataLoaderIterator(*dataset_, batch_size_, max_batch_idx_, max_batch_idx_); + return DataLoaderIterator(*dataset_, batch_size_, num_global_batches_, num_global_batches_); } +size_t DataLoader::NumGlobalBatches() const { return num_global_batches_; } + DistributedDataLoader::DistributedDataLoader(const std::shared_ptr &dataset, size_t batch_size, size_t ddp_rank, size_t ddp_world_size) - : DataLoader(dataset, batch_size), ddp_rank_(ddp_rank), ddp_world_size_(ddp_world_size) {} + : DataLoader(dataset, batch_size), ddp_rank_(ddp_rank), ddp_world_size_(ddp_world_size) { + CHECK_GT(ddp_world_size_, 0); + CHECK_LT(ddp_rank_, ddp_world_size_); + const size_t global_batch_size = ddp_world_size_ * batch_size_; + CHECK_GE(dataset_->Size(), global_batch_size) + << "DistributedDataLoader needs at least one full global batch. dataset_size=" << dataset_->Size() + << ", global_batch_size=" << global_batch_size << " (" << batch_size_ << " per rank * " << ddp_world_size_ + << " ranks). Reduce batch size/world size or use a larger dataset."; + num_global_batches_ = dataset_->Size() / global_batch_size; +} DataLoaderIterator DistributedDataLoader::begin() const { - return DataLoaderIterator(*dataset_, batch_size_, ddp_rank_, max_batch_idx_, ddp_rank_, ddp_world_size_); + return DataLoaderIterator(*dataset_, batch_size_, 0, num_global_batches_, ddp_rank_, ddp_world_size_); } DataLoaderIterator DistributedDataLoader::end() const { - return DataLoaderIterator(*dataset_, batch_size_, max_batch_idx_, max_batch_idx_, ddp_rank_, ddp_world_size_); + return DataLoaderIterator(*dataset_, batch_size_, num_global_batches_, num_global_batches_, ddp_rank_, + ddp_world_size_); } } // namespace infini_train diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 96776585..a26d27d8 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -7,6 +7,9 @@ include(${CMAKE_SOURCE_DIR}/cmake/test_macros.cmake) # Common test utilities add_subdirectory(common) +# DataLoader tests +add_subdirectory(dataloader) + # Tensor tests add_subdirectory(tensor) diff --git a/tests/dataloader/CMakeLists.txt b/tests/dataloader/CMakeLists.txt new file mode 100644 index 00000000..7dcf5c88 --- /dev/null +++ b/tests/dataloader/CMakeLists.txt @@ -0,0 +1,8 @@ +# ========================================================================== +# DataLoader tests +# ========================================================================== + +infini_train_add_test(test_dataloader + SOURCES test_dataloader.cc + LABELS cpu +) diff --git a/tests/dataloader/test_dataloader.cc b/tests/dataloader/test_dataloader.cc new file mode 100644 index 00000000..837232d1 --- /dev/null +++ b/tests/dataloader/test_dataloader.cc @@ -0,0 +1,101 @@ +#include +#include +#include +#include + +#include "gtest/gtest.h" + +#include "infini_train/include/dataloader.h" +#include "infini_train/include/dataset.h" +#include "infini_train/include/tensor.h" + +using namespace infini_train; + +namespace { +class IndexDataset : public Dataset { +public: + explicit IndexDataset(size_t size) : size_(size) {} + + std::pair, std::shared_ptr> operator[](size_t idx) const override { + auto data = std::make_shared(std::vector{1}, DataType::kINT64); + auto label = std::make_shared(std::vector{1}, DataType::kINT64); + *static_cast(data->DataPtr()) = static_cast(idx); + *static_cast(label->DataPtr()) = static_cast(idx + 1000); + return {data, label}; + } + + size_t Size() const override { return size_; } + +private: + size_t size_ = 0; +}; + +std::vector TensorValues(const std::shared_ptr &tensor) { + const auto *data = static_cast(tensor->DataPtr()); + return std::vector(data, data + tensor->NumElements()); +} +} // namespace + +TEST(DataLoaderTest, RegularDataLoaderKeepsPartialLastBatch) { + DataLoader loader(std::make_shared(5), 2); + + std::vector> batches; + for (const auto &[x, y] : loader) { batches.push_back(TensorValues(x)); } + + ASSERT_EQ(batches.size(), 3); + EXPECT_EQ(batches[0], (std::vector{0, 1})); + EXPECT_EQ(batches[1], (std::vector{2, 3})); + EXPECT_EQ(batches[2], (std::vector{4})); +} + +TEST(DataLoaderTest, DistributedDataLoaderSlicesFullGlobalBatchesByRank) { + const auto dataset = std::make_shared(13); + const size_t batch_size = 2; + const size_t world_size = 3; + + DistributedDataLoader rank0(dataset, batch_size, 0, world_size); + DistributedDataLoader rank1(dataset, batch_size, 1, world_size); + DistributedDataLoader rank2(dataset, batch_size, 2, world_size); + + auto r0 = rank0.begin(); + auto r1 = rank1.begin(); + auto r2 = rank2.begin(); + + EXPECT_EQ(TensorValues((*r0).first), (std::vector{0, 1})); + EXPECT_EQ(TensorValues((*r1).first), (std::vector{2, 3})); + EXPECT_EQ(TensorValues((*r2).first), (std::vector{4, 5})); + + ++r0; + ++r1; + ++r2; + + EXPECT_EQ(TensorValues((*r0).first), (std::vector{6, 7})); + EXPECT_EQ(TensorValues((*r1).first), (std::vector{8, 9})); + EXPECT_EQ(TensorValues((*r2).first), (std::vector{10, 11})); + + ++r0; + ++r1; + ++r2; + + EXPECT_EQ(r0, rank0.end()); + EXPECT_EQ(r1, rank1.end()); + EXPECT_EQ(r2, rank2.end()); +} + +TEST(DataLoaderTest, DistributedDataLoaderCyclesOnGlobalBatchBoundary) { + DistributedDataLoader loader(std::make_shared(13), 2, 2, 3); + auto iter = loader.begin(); + + EXPECT_EQ(loader.NumGlobalBatches(), 2); + + EXPECT_EQ(TensorValues((*iter).first), (std::vector{4, 5})); + ++iter; + EXPECT_EQ(iter.GlobalBatchIndex(), 1); + + EXPECT_EQ(TensorValues((*iter).first), (std::vector{10, 11})); + ++iter; + EXPECT_EQ(iter, loader.end()); + + iter = loader.begin(); + EXPECT_EQ(TensorValues((*iter).first), (std::vector{4, 5})); +}