diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 2551880e..023469cb 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -383,13 +383,20 @@ void Train(const nn::parallel::Rank &rank) { start_step = resume_result.global_step; size_t consumed_train_samples = resume_result.consumed_train_samples; - // TODO(jym): Replace with Sampler abstraction when available. - // Skip dataloader to resume from the correct batch position. - if (consumed_train_samples > 0) { - const size_t num_skips - = DataLoaderBatchesToSkip(consumed_train_samples, train_loader_batch_size, ddp_world_size); - for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } - } + const size_t consumed_dataloader_steps + = DataLoaderStepsToSkip(consumed_train_samples, train_loader_batch_size, ddp_world_size); + train_iter.SeekDataLoaderStep(consumed_dataloader_steps % train_loader.NumDataLoaderSteps()); + 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_train_samples += train_loader_batch_size * ddp_world_size; + return batch; + }; auto save_checkpoint = [&](const std::filesystem::path &save_dir, int64_t global_step) { SaveCheckpoint({ @@ -463,11 +470,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_train_samples += static_cast(FLAGS_batch_size) * ddp_world_size; + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -499,11 +502,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_train_samples += train_loader_batch_size * ddp_world_size; + 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 ccfca86a..09086ef2 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -365,13 +365,20 @@ void Train(const nn::parallel::Rank &rank) { start_step = resume_result.global_step; size_t consumed_train_samples = resume_result.consumed_train_samples; - // TODO(jym): Replace with Sampler abstraction when available. - // Skip dataloader to resume from the correct batch position. - if (consumed_train_samples > 0) { - const size_t num_skips - = DataLoaderBatchesToSkip(consumed_train_samples, train_loader_batch_size, ddp_world_size); - for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } - } + const size_t consumed_dataloader_steps + = DataLoaderStepsToSkip(consumed_train_samples, train_loader_batch_size, ddp_world_size); + train_iter.SeekDataLoaderStep(consumed_dataloader_steps % train_loader.NumDataLoaderSteps()); + 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_train_samples += train_loader_batch_size * ddp_world_size; + return batch; + }; auto save_checkpoint = [&](const std::filesystem::path &save_dir, int64_t global_step) { SaveCheckpoint({ @@ -443,11 +450,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_train_samples += static_cast(FLAGS_batch_size) * ddp_world_size; + auto [x, y] = next_train_batch(); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -478,11 +481,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_train_samples += train_loader_batch_size * ddp_world_size; + 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/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 13490ccc..088463dd 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -62,4 +62,4 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & void SaveCheckpoint(const SaveCheckpointArgs &args); -size_t DataLoaderBatchesToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size); +size_t DataLoaderStepsToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size); diff --git a/infini_train/include/dataloader.h b/infini_train/include/dataloader.h index ad7fbcda..a75a719c 100644 --- a/infini_train/include/dataloader.h +++ b/infini_train/include/dataloader.h @@ -12,9 +12,6 @@ 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, - size_t ddp_rank = 0, size_t ddp_world_size = 1); - std::pair, std::shared_ptr> operator*() const; DataLoaderIterator &operator++(); @@ -24,11 +21,20 @@ class DataLoaderIterator { friend bool operator!=(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs); friend bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs); + size_t DataLoaderStep() const; + DataLoaderIterator &SeekDataLoaderStep(size_t dataloader_step); + private: + friend class DataLoader; + friend class DistributedDataLoader; + + DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t dataloader_step, size_t num_dataloader_steps, + size_t ddp_rank, size_t ddp_world_size); + const Dataset *dataset_ = nullptr; // not owned size_t batch_size_ = 0; - size_t batch_idx_ = 0; - size_t max_batch_idx_ = 0; + size_t dataloader_step_ = 0; + size_t num_dataloader_steps_ = 0; size_t ddp_rank_ = 0; size_t ddp_world_size_ = 1; }; @@ -40,10 +46,12 @@ class DataLoader { virtual DataLoaderIterator begin() const; virtual DataLoaderIterator end() const; + size_t NumDataLoaderSteps() const; + protected: std::shared_ptr dataset_; size_t batch_size_ = 0; - size_t max_batch_idx_ = 0; + size_t num_dataloader_steps_ = 0; }; class DistributedDataLoader : public DataLoader { diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index c6e31cdd..c946f475 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -120,15 +120,15 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { } } -size_t DataLoaderBatchesToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size) { +size_t DataLoaderStepsToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size) { CHECK_GT(local_batch_size, 0); CHECK_GT(ddp_world_size, 0); CHECK_LE(local_batch_size, std::numeric_limits::max() / ddp_world_size) << "Data loader batch size overflows size_t"; - const size_t global_loader_batch_size = local_batch_size * ddp_world_size; - CHECK_EQ(consumed_train_samples % global_loader_batch_size, 0) + const size_t samples_per_dataloader_step = local_batch_size * ddp_world_size; + CHECK_EQ(consumed_train_samples % samples_per_dataloader_step, 0) << "consumed_train_samples=" << consumed_train_samples << " does not align with current local_batch_size=" << local_batch_size << " and ddp_world_size=" << ddp_world_size; - return consumed_train_samples / global_loader_batch_size; + return consumed_train_samples / samples_per_dataloader_step; } diff --git a/infini_train/src/dataloader.cc b/infini_train/src/dataloader.cc index 322df553..3c252884 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 dataloader_step, + size_t num_dataloader_steps, size_t ddp_rank, size_t ddp_world_size) + : dataset_(&dataset), batch_size_(batch_size), dataloader_step_(dataloader_step), + num_dataloader_steps_(num_dataloader_steps), 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 + dataloader_step */ 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(dataloader_step_, num_dataloader_steps_) + << "Cannot dereference DataLoader end iterator. dataloader_step=" << dataloader_step_ + << ", num_dataloader_steps=" << num_dataloader_steps_ << ", ddp_rank=" << ddp_rank_ + << ", ddp_world_size=" << ddp_world_size_; + const size_t start_idx = (dataloader_step_ * ddp_world_size_ + ddp_rank_) * batch_size_; + CHECK_LT(start_idx, dataset_->Size()) + << "DataLoader batch starts past dataset end. dataloader_step=" << dataloader_step_ + << ", 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_); + dataloader_step_ = std::min(dataloader_step_ + 1, num_dataloader_steps_); return *this; } @@ -68,34 +85,64 @@ 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.dataloader_step_ < rhs.dataloader_step_; +} bool operator!=(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { - return lhs.batch_idx_ != rhs.batch_idx_; + return lhs.dataloader_step_ != rhs.dataloader_step_; } bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) { - return lhs.batch_idx_ == rhs.batch_idx_; + return lhs.dataloader_step_ == rhs.dataloader_step_; +} + +size_t DataLoaderIterator::DataLoaderStep() const { return dataloader_step_; } + +DataLoaderIterator &DataLoaderIterator::SeekDataLoaderStep(size_t dataloader_step) { + CHECK_LE(dataloader_step, num_dataloader_steps_) + << "Cannot seek past DataLoader end. dataloader_step=" << dataloader_step + << ", num_dataloader_steps=" << num_dataloader_steps_; + dataloader_step_ = dataloader_step; + 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) { + CHECK(dataset_ != nullptr) << "DataLoader dataset must not be null"; + CHECK_GT(batch_size_, 0) << "DataLoader batch_size must be greater than zero"; + num_dataloader_steps_ = 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_dataloader_steps_, 0, 1); +} DataLoaderIterator DataLoader::end() const { - return DataLoaderIterator(*dataset_, batch_size_, max_batch_idx_, max_batch_idx_); + return DataLoaderIterator(*dataset_, batch_size_, num_dataloader_steps_, num_dataloader_steps_, 0, 1); } +size_t DataLoader::NumDataLoaderSteps() const { return num_dataloader_steps_; } + 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 samples_per_dataloader_step = ddp_world_size_ * batch_size_; + CHECK_GE(dataset_->Size(), samples_per_dataloader_step) + << "DistributedDataLoader needs enough samples for one DataLoader step. dataset_size=" << dataset_->Size() + << ", samples_per_dataloader_step=" << samples_per_dataloader_step << " (" << batch_size_ << " per rank * " + << ddp_world_size_ << " ranks). Reduce batch size/world size or use a larger dataset."; + num_dataloader_steps_ = dataset_->Size() / samples_per_dataloader_step; +} 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_dataloader_steps_, 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_dataloader_steps_, num_dataloader_steps_, ddp_rank_, + ddp_world_size_); } } // namespace infini_train diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 3bfaa548..43bef143 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -12,6 +12,9 @@ add_subdirectory(distributed) # Module tests add_subdirectory(module) +# DataLoader tests +add_subdirectory(dataloader) + # Tensor tests add_subdirectory(tensor) diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index ec4d61e8..d5c8085e 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -105,18 +105,18 @@ TEST_P(TrainerStateTest, RoundTrip) { std::filesystem::remove_all(dir); } -TEST_P(TrainerStateTest, DataLoaderSkipUsesCurrentBatchConfiguration) { - EXPECT_EQ(DataLoaderBatchesToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/4, - /*ddp_world_size=*/2), +TEST_P(TrainerStateTest, DataLoaderStepsToSkipUsesCurrentBatchConfiguration) { + EXPECT_EQ(DataLoaderStepsToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/4, + /*ddp_world_size=*/2), 50); - EXPECT_EQ(DataLoaderBatchesToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/8, - /*ddp_world_size=*/2), + EXPECT_EQ(DataLoaderStepsToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/8, + /*ddp_world_size=*/2), 25); } -TEST_P(TrainerStateTest, DataLoaderSkipRejectsUnalignedConfiguration) { - EXPECT_DEATH(DataLoaderBatchesToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/6, - /*ddp_world_size=*/4), +TEST_P(TrainerStateTest, DataLoaderStepsToSkipRejectsUnalignedConfiguration) { + EXPECT_DEATH(DataLoaderStepsToSkip(/*consumed_train_samples=*/400, /*local_batch_size=*/6, + /*ddp_world_size=*/4), "does not align"); } 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..ca5d400b --- /dev/null +++ b/tests/dataloader/test_dataloader.cc @@ -0,0 +1,128 @@ +#include +#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; + +static_assert(!std::is_constructible_v); + +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, RejectsInvalidConstructorArguments) { + EXPECT_DEATH({ DataLoader loader(nullptr, 2); }, "dataset must not be null"); + EXPECT_DEATH({ DataLoader loader(std::make_shared(5), 0); }, "batch_size must be greater than zero"); +} + +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})); + + auto iter = loader.begin(); + iter.SeekDataLoaderStep(2); + EXPECT_EQ(TensorValues((*iter).first), (std::vector{4})); +} + +TEST(DataLoaderTest, DistributedDataLoaderPartitionsEachStepByRank) { + 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, SeekDataLoaderStepSupportsResumeAndEnd) { + DistributedDataLoader loader(std::make_shared(13), 2, 2, 3); + auto iter = loader.begin(); + + EXPECT_EQ(loader.NumDataLoaderSteps(), 2); + + iter.SeekDataLoaderStep(1); + EXPECT_EQ(iter.DataLoaderStep(), 1); + EXPECT_EQ(TensorValues((*iter).first), (std::vector{10, 11})); + + iter.SeekDataLoaderStep(loader.NumDataLoaderSteps()); + EXPECT_EQ(iter, loader.end()); + EXPECT_DEATH(iter.SeekDataLoaderStep(loader.NumDataLoaderSteps() + 1), "Cannot seek past DataLoader end"); + + size_t consumed_dataloader_steps = 5; + iter = loader.begin(); + iter.SeekDataLoaderStep(consumed_dataloader_steps % loader.NumDataLoaderSteps()); + auto next_batch = [&]() { + auto batch = *iter; + ++iter; + if (iter == loader.end()) { + iter = loader.begin(); + } + ++consumed_dataloader_steps; + return batch; + }; + + EXPECT_EQ(TensorValues(next_batch().first), (std::vector{10, 11})); + EXPECT_EQ(iter.DataLoaderStep(), 0); + EXPECT_EQ(TensorValues(next_batch().first), (std::vector{4, 5})); + EXPECT_EQ(consumed_dataloader_steps, 7); +}