Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 20 additions & 13 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -383,14 +383,29 @@ 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.
auto advance_train_iter = [&]() {
++train_iter;
if (train_iter == train_loader.end()) {
train_iter = train_loader.begin();
}
};

// TODO(jym): Move resume position handling into a Sampler abstraction when available.
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; }
for (size_t i = 0; i < num_skips; ++i) { advance_train_iter(); }
}

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
advance_train_iter();
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({
.save_dir = save_dir,
Expand Down Expand Up @@ -463,11 +478,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<size_t>(FLAGS_batch_size) * ddp_world_size;
auto [x, y] = next_train_batch();
x = std::make_shared<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

Expand Down Expand Up @@ -499,11 +510,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<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

Expand Down
33 changes: 20 additions & 13 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -365,14 +365,29 @@ 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.
auto advance_train_iter = [&]() {
++train_iter;
if (train_iter == train_loader.end()) {
train_iter = train_loader.begin();
}
};

// TODO(jym): Move resume position handling into a Sampler abstraction when available.
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; }
for (size_t i = 0; i < num_skips; ++i) { advance_train_iter(); }
}

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
advance_train_iter();
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({
.save_dir = save_dir,
Expand Down Expand Up @@ -443,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_train_samples += static_cast<size_t>(FLAGS_batch_size) * ddp_world_size;
auto [x, y] = next_train_batch();
x = std::make_shared<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

Expand Down Expand Up @@ -478,11 +489,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<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

Expand Down
13 changes: 8 additions & 5 deletions infini_train/include/dataloader.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<Tensor>, std::shared_ptr<Tensor>> operator*() const;

DataLoaderIterator &operator++();
Expand All @@ -25,10 +22,16 @@ class DataLoaderIterator {
friend bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs);

private:
friend class DataLoader;
friend class DistributedDataLoader;

DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t batch_idx, size_t num_batches, 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 num_batches_ = 0;
size_t ddp_rank_ = 0;
size_t ddp_world_size_ = 1;
};
Expand All @@ -43,7 +46,7 @@ class DataLoader {
protected:
std::shared_ptr<Dataset> dataset_;
size_t batch_size_ = 0;
size_t max_batch_idx_ = 0;
size_t num_batches_ = 0;
};

class DistributedDataLoader : public DataLoader {
Expand Down
53 changes: 42 additions & 11 deletions infini_train/src/dataloader.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <algorithm>
#include <cstddef>
#include <cstring>
#include <functional>
#include <numeric>
#include <utility>

Expand All @@ -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<Tensor> Stack(const std::vector<std::shared_ptr<Tensor>> &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<int64_t>());
Expand All @@ -35,9 +42,9 @@ std::shared_ptr<Tensor> Stack(const std::vector<std::shared_ptr<Tensor>> &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),
DataLoaderIterator::DataLoaderIterator(const Dataset &dataset, size_t batch_size, size_t batch_idx, size_t num_batches,
size_t ddp_rank, size_t ddp_world_size)
: dataset_(&dataset), batch_size_(batch_size), batch_idx_(batch_idx), num_batches_(num_batches),
ddp_rank_(ddp_rank), ddp_world_size_(ddp_world_size){};

std::pair<std::shared_ptr<Tensor>, std::shared_ptr<Tensor>> DataLoaderIterator::operator*() const {
Expand All @@ -49,7 +56,16 @@ std::pair<std::shared_ptr<Tensor>, std::shared_ptr<Tensor>> DataLoaderIterator::
*/
std::vector<std::shared_ptr<Tensor>> data_vec;
std::vector<std::shared_ptr<Tensor>> label_vec;
for (int idx = batch_idx_ * batch_size_; idx < (batch_idx_ + 1) * batch_size_ && idx < dataset_->Size(); ++idx) {
CHECK_LT(batch_idx_, num_batches_) << "Cannot dereference DataLoader end iterator. batch_idx=" << batch_idx_
<< ", num_batches=" << num_batches_ << ", ddp_rank=" << ddp_rank_
<< ", ddp_world_size=" << ddp_world_size_;
const size_t start_idx = (batch_idx_ * ddp_world_size_ + ddp_rank_) * batch_size_;
CHECK_LT(start_idx, dataset_->Size())
<< "DataLoader batch starts past dataset end. batch_idx=" << 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));
Expand All @@ -58,7 +74,7 @@ std::pair<std::shared_ptr<Tensor>, std::shared_ptr<Tensor>> DataLoaderIterator::
}

DataLoaderIterator &DataLoaderIterator::operator++() {
batch_idx_ = std::min(batch_idx_ + ddp_world_size_, max_batch_idx_);
batch_idx_ = std::min(batch_idx_ + 1, num_batches_);
return *this;
}

Expand All @@ -79,23 +95,38 @@ bool operator==(const DataLoaderIterator &lhs, const DataLoaderIterator &rhs) {
}

DataLoader::DataLoader(const std::shared_ptr<Dataset> &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_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_batches_, 0, 1);
}

DataLoaderIterator DataLoader::end() const {
return DataLoaderIterator(*dataset_, batch_size_, max_batch_idx_, max_batch_idx_);
return DataLoaderIterator(*dataset_, batch_size_, num_batches_, num_batches_, 0, 1);
}

DistributedDataLoader::DistributedDataLoader(const std::shared_ptr<Dataset> &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 enough samples for one 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_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_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_batches_, num_batches_, ddp_rank_, ddp_world_size_);
}
} // namespace infini_train
3 changes: 3 additions & 0 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@ add_subdirectory(distributed)
# Module tests
add_subdirectory(module)

# DataLoader tests
add_subdirectory(dataloader)

# Tensor tests
add_subdirectory(tensor)

Expand Down
8 changes: 8 additions & 0 deletions tests/dataloader/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
# ==========================================================================
# DataLoader tests
# ==========================================================================

infini_train_add_test(test_dataloader
SOURCES test_dataloader.cc
LABELS cpu
)
104 changes: 104 additions & 0 deletions tests/dataloader/test_dataloader.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
#include <cstdint>
#include <memory>
#include <type_traits>
#include <utility>
#include <vector>

#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<DataLoaderIterator, const Dataset &, size_t, size_t, size_t, size_t, size_t>);

namespace {
class IndexDataset : public Dataset {
public:
explicit IndexDataset(size_t size) : size_(size) {}

std::pair<std::shared_ptr<Tensor>, std::shared_ptr<Tensor>> operator[](size_t idx) const override {
auto data = std::make_shared<Tensor>(std::vector<int64_t>{1}, DataType::kINT64);
auto label = std::make_shared<Tensor>(std::vector<int64_t>{1}, DataType::kINT64);
*static_cast<int64_t *>(data->DataPtr()) = static_cast<int64_t>(idx);
*static_cast<int64_t *>(label->DataPtr()) = static_cast<int64_t>(idx + 1000);
return {data, label};
}

size_t Size() const override { return size_; }

private:
size_t size_ = 0;
};

std::vector<int64_t> TensorValues(const std::shared_ptr<Tensor> &tensor) {
const auto *data = static_cast<const int64_t *>(tensor->DataPtr());
return std::vector<int64_t>(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<IndexDataset>(5), 0); }, "batch_size must be greater than zero");
}

TEST(DataLoaderTest, RegularDataLoaderKeepsPartialLastBatch) {
DataLoader loader(std::make_shared<IndexDataset>(5), 2);

std::vector<std::vector<int64_t>> batches;
for (const auto &[x, y] : loader) { batches.push_back(TensorValues(x)); }

ASSERT_EQ(batches.size(), 3);
EXPECT_EQ(batches[0], (std::vector<int64_t>{0, 1}));
EXPECT_EQ(batches[1], (std::vector<int64_t>{2, 3}));
EXPECT_EQ(batches[2], (std::vector<int64_t>{4}));
}

TEST(DataLoaderTest, DistributedDataLoaderPartitionsEachStepByRank) {
const auto dataset = std::make_shared<IndexDataset>(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<int64_t>{0, 1}));
EXPECT_EQ(TensorValues((*r1).first), (std::vector<int64_t>{2, 3}));
EXPECT_EQ(TensorValues((*r2).first), (std::vector<int64_t>{4, 5}));

++r0;
++r1;
++r2;

EXPECT_EQ(TensorValues((*r0).first), (std::vector<int64_t>{6, 7}));
EXPECT_EQ(TensorValues((*r1).first), (std::vector<int64_t>{8, 9}));
EXPECT_EQ(TensorValues((*r2).first), (std::vector<int64_t>{10, 11}));

++r0;
++r1;
++r2;

EXPECT_EQ(r0, rank0.end());
EXPECT_EQ(r1, rank1.end());
EXPECT_EQ(r2, rank2.end());
}

TEST(DataLoaderTest, DistributedDataLoaderCanRestartAtGlobalBatchBoundary) {
DistributedDataLoader loader(std::make_shared<IndexDataset>(13), 2, 2, 3);
auto iter = loader.begin();

++iter;
EXPECT_EQ(TensorValues((*iter).first), (std::vector<int64_t>{10, 11}));

++iter;
EXPECT_EQ(iter, loader.end());
iter = loader.begin();
EXPECT_EQ(TensorValues((*iter).first), (std::vector<int64_t>{4, 5}));
}
Loading