Skip to content
Open
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
34 changes: 15 additions & 19 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand Down Expand Up @@ -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<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

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

Expand Down
34 changes: 15 additions & 19 deletions example/llama3/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand Down Expand Up @@ -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<Tensor>(x->To(device));
y = std::make_shared<Tensor>(y->To(device));

Expand Down Expand Up @@ -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<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,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<Tensor>, std::shared_ptr<Tensor>> operator*() const;
Expand All @@ -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;
};
Expand All @@ -42,10 +43,12 @@ class DataLoader {
virtual DataLoaderIterator begin() const;
virtual DataLoaderIterator end() const;

size_t NumGlobalBatches() const;

protected:
std::shared_ptr<Dataset> dataset_;
size_t batch_size_ = 0;
size_t max_batch_idx_ = 0;
size_t num_global_batches_ = 0;
};

class DistributedDataLoader : public DataLoader {
Expand Down
75 changes: 58 additions & 17 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,21 +42,31 @@ 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),
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<Tensor>, std::shared_ptr<Tensor>> 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<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(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));
Expand All @@ -58,7 +75,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_);
global_batch_idx_ = std::min(global_batch_idx_ + 1, num_global_batches_);
return *this;
}

Expand All @@ -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> &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> &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
3 changes: 3 additions & 0 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)

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
)
Loading
Loading