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

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

Expand Down
2 changes: 1 addition & 1 deletion infini_train/include/checkpoint/checkpoint_manager.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);
20 changes: 14 additions & 6 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 @@ -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);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

“从第 xx 个 step 继续读取” 的状态控制应当是 Sampler 的职责:Sampler 根据恢复进度决定起始 sample/index,DataLoader 本身不需要额外维护一份可 seek 的状态。

Megatron 也是在 Sampler 层通过 consumed_samples 控制数据恢复位置:
https://github.com/NVIDIA/Megatron-LM/blob/f7f584d7a04c1891bc583a1213ebb1dbcac4e158/megatron/training/datasets/data_samplers.py#L172

建议暂时不要在 DataLoader / DataLoaderIterator 中引入额外的 dataloader_step 和 SeekDataLoaderStep 状态。现阶段继续在 main 侧通过迭代并跳过前 N 个 batch 的方式完成断点恢复,作为临时方案;同时保留 Sampler 的 TODO,后续引入 Sampler 抽象后,再将恢复位置的控制下沉到 Sampler。


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;
};
Expand All @@ -40,10 +46,12 @@ class DataLoader {
virtual DataLoaderIterator begin() const;
virtual DataLoaderIterator end() const;

size_t NumDataLoaderSteps() const;

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

class DistributedDataLoader : public DataLoader {
Expand Down
8 changes: 4 additions & 4 deletions infini_train/src/checkpoint/checkpoint_manager.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<size_t>::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;
}
79 changes: 63 additions & 16 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 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<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
dataloader_step
*/
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(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));
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_);
dataloader_step_ = std::min(dataloader_step_ + 1, num_dataloader_steps_);
return *this;
}

Expand All @@ -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> &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> &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
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
16 changes: 8 additions & 8 deletions tests/checkpoint/test_trainer_state.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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");
}

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