From d98e7006a2a408f219a870037d69e576e5977df5 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Tue, 4 Aug 2026 07:57:31 +0000 Subject: [PATCH 1/2] feat: track checkpoint progress in micro-batches --- example/gpt2/main.cc | 22 +++++++++++-------- example/llama3/main.cc | 22 +++++++++++-------- infini_train/include/checkpoint/checkpoint.h | 2 +- .../include/checkpoint/checkpoint_manager.h | 4 ++-- infini_train/src/checkpoint/checkpoint.cc | 8 +++---- .../src/checkpoint/checkpoint_manager.cc | 8 +++---- .../test_checkpoint_serialization.cc | 4 ++-- tests/checkpoint/test_lr_scheduler_state.cc | 4 ++-- tests/checkpoint/test_trainer_state.cc | 10 ++++----- 9 files changed, 46 insertions(+), 38 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 976a7031..1a77a01b 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -382,15 +382,19 @@ void Train(const nn::parallel::Rank &rank) { .state = state, .lr_scheduler = scheduler}); start_step = resume_result.global_step; - size_t consumed_batches = resume_result.consumed_batches; + size_t consumed_micro_batches = resume_result.consumed_micro_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; + if (consumed_micro_batches > 0) { + const size_t start = train_iter.BatchIndex(); + CHECK(pp_world_size == 1 || consumed_micro_batches % num_micro_batches == 0); + const size_t consumed_loader_batches + = pp_world_size > 1 ? consumed_micro_batches / num_micro_batches : consumed_micro_batches; + const size_t target = consumed_loader_batches + static_cast(ddp_rank); + CHECK_GE(target, start); + CHECK_EQ((target - start) % ddp_world_size, 0); + const size_t num_skips = (target - start) / ddp_world_size; for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } } @@ -398,7 +402,7 @@ void Train(const nn::parallel::Rank &rank) { SaveCheckpoint({ .save_dir = save_dir, .global_step = global_step, - .consumed_batches = consumed_batches, + .consumed_micro_batches = consumed_micro_batches, .n_layer = model_config.n_layer, .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, @@ -470,7 +474,7 @@ void Train(const nn::parallel::Rank &rank) { // 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(); + consumed_micro_batches = train_iter.BatchIndex() - static_cast(ddp_rank); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -504,7 +508,7 @@ void Train(const nn::parallel::Rank &rank) { // 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(); + consumed_micro_batches = (train_iter.BatchIndex() - static_cast(ddp_rank)) * num_micro_batches; 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 2a209614..a20a034e 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -364,15 +364,19 @@ void Train(const nn::parallel::Rank &rank) { .lr_scheduler = scheduler}); start_step = resume_result.global_step; - size_t consumed_batches = resume_result.consumed_batches; + size_t consumed_micro_batches = resume_result.consumed_micro_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; + if (consumed_micro_batches > 0) { + const size_t start = train_iter.BatchIndex(); + CHECK(pp_world_size == 1 || consumed_micro_batches % num_micro_batches == 0); + const size_t consumed_loader_batches + = pp_world_size > 1 ? consumed_micro_batches / num_micro_batches : consumed_micro_batches; + const size_t target = consumed_loader_batches + static_cast(ddp_rank); + CHECK_GE(target, start); + CHECK_EQ((target - start) % ddp_world_size, 0); + const size_t num_skips = (target - start) / ddp_world_size; for (size_t i = 0; i < num_skips; ++i) { ++train_iter; } } @@ -380,7 +384,7 @@ void Train(const nn::parallel::Rank &rank) { SaveCheckpoint({ .save_dir = save_dir, .global_step = global_step, - .consumed_batches = consumed_batches, + .consumed_micro_batches = consumed_micro_batches, .n_layer = model_config.n_layer, .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, @@ -450,7 +454,7 @@ void Train(const nn::parallel::Rank &rank) { // 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(); + consumed_micro_batches = train_iter.BatchIndex() - static_cast(ddp_rank); x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -483,7 +487,7 @@ void Train(const nn::parallel::Rank &rank) { // 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(); + consumed_micro_batches = (train_iter.BatchIndex() - static_cast(ddp_rank)) * num_micro_batches; x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index 0d3cee6b..ecb5ce16 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -17,7 +17,7 @@ class Module; struct TrainerState { int64_t global_step = 0; - int64_t consumed_batches = 0; + int64_t consumed_micro_batches = 0; int64_t n_layer = 0; int64_t n_head = 0; int64_t n_kv_head = 0; diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 7e7b5eb7..92a54632 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -34,13 +34,13 @@ struct ResumeFromCheckpointArgs { struct ResumeFromCheckpointResult { int global_step = 0; - size_t consumed_batches = 0; + size_t consumed_micro_batches = 0; }; struct SaveCheckpointArgs { std::filesystem::path save_dir; int64_t global_step = 0; - size_t consumed_batches = 0; + size_t consumed_micro_batches = 0; int64_t n_layer = 0; int64_t n_head = 0; int64_t n_kv_head = 0; diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index 988b5120..38cb5c19 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -237,8 +237,8 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m } LOG(ERROR) << "[CKPT] Load done: global_step=" << state.global_step - << ", consumed_batches =" << state.consumed_batches << ", topology(ddp,tp,sp,pp)=(" << state.ddp_size - << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; + << ", consumed_micro_batches =" << state.consumed_micro_batches << ", topology(ddp,tp,sp,pp)=(" + << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; } void Checkpoint::SaveStateDict(const std::filesystem::path &path, @@ -320,7 +320,7 @@ void Checkpoint::SaveTrainerState(const std::filesystem::path &path, const Train ofs << " \"n_embd\": " << state.n_embd << ",\n"; ofs << " \"vocab_size\": " << state.vocab_size << ",\n"; ofs << " \"global_step\": " << state.global_step << ",\n"; - ofs << " \"consumed_batches\": " << state.consumed_batches << ",\n"; + ofs << " \"consumed_micro_batches\": " << state.consumed_micro_batches << ",\n"; ofs << " \"ddp_size\": " << state.ddp_size << ",\n"; ofs << " \"tp_size\": " << state.tp_size << ",\n"; ofs << " \"sp_size\": " << state.sp_size << ",\n"; @@ -341,7 +341,7 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.n_embd = ExtractNumberField(content, "n_embd", 0); state.vocab_size = ExtractNumberField(content, "vocab_size", 0); state.global_step = ExtractNumberField(content, "global_step", 0); - state.consumed_batches = ExtractNumberField(content, "consumed_batches", 0); + state.consumed_micro_batches = ExtractNumberField(content, "consumed_micro_batches", 0); state.ddp_size = ExtractNumberField(content, "ddp_size", 1); state.tp_size = ExtractNumberField(content, "tp_size", 1); state.sp_size = ExtractNumberField(content, "sp_size", 1); diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index 51c20f12..427ec2f4 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -63,10 +63,10 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & CHECK_EQ(args.state.pp_size, pp_world_size) << "PP size mismatch: checkpoint has PP=" << args.state.pp_size << ", but current run has PP=" << pp_world_size; - result.consumed_batches = static_cast(std::max(args.state.consumed_batches, 0)); + result.consumed_micro_batches = static_cast(std::max(args.state.consumed_micro_batches, 0)); if (args.rank.IsMainRank()) { - LOG(INFO) << std::format("Resume training from step {}, consumed_batches {}", args.state.global_step, - args.state.consumed_batches); + LOG(INFO) << std::format("Resume training from step {}, consumed_micro_batches {}", args.state.global_step, + args.state.consumed_micro_batches); } return result; @@ -77,7 +77,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { TrainerState state; state.global_step = args.global_step; - state.consumed_batches = static_cast(args.consumed_batches); + state.consumed_micro_batches = static_cast(args.consumed_micro_batches); state.n_layer = args.n_layer; state.n_head = args.n_head; state.n_kv_head = args.n_kv_head; diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 495dcf27..c03b8462 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -29,7 +29,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { *model1->mutable_parameter("bias") = p2; auto opt1 = std::make_shared(model1->Parameters(), 0.01); - TrainerState saved{.global_step = 42, .consumed_batches = 100}; + TrainerState saved{.global_step = 42, .consumed_micro_batches = 100}; Checkpoint::Save(dir, *model1, opt1.get(), saved, nullptr); auto model2 = std::make_shared(3, 2, true, GetDevice()); @@ -45,7 +45,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { Checkpoint::Load(dir, *model2, opt2.get(), loaded, nullptr); EXPECT_EQ(loaded.global_step, 42); - EXPECT_EQ(loaded.consumed_batches, 100); + EXPECT_EQ(loaded.consumed_micro_batches, 100); auto w1_cpu = model2->parameter("weight")->To(Device()); const float *data = static_cast(w1_cpu.DataPtr()); diff --git a/tests/checkpoint/test_lr_scheduler_state.cc b/tests/checkpoint/test_lr_scheduler_state.cc index fc49fb3d..9df7238b 100644 --- a/tests/checkpoint/test_lr_scheduler_state.cc +++ b/tests/checkpoint/test_lr_scheduler_state.cc @@ -59,7 +59,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { auto sched1 = CreateLRScheduler(opt1, MakeSchedulerConfig()); StepTimes(sched1, 3); - TrainerState saved{.global_step = 3, .consumed_batches = 12}; + TrainerState saved{.global_step = 3, .consumed_micro_batches = 12}; Checkpoint::Save(dir, *model1, nullptr, saved, sched1.get()); EXPECT_TRUE(std::filesystem::exists(dir / "lr_scheduler.ckpt")); @@ -71,7 +71,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { Checkpoint::Load(dir, *model2, nullptr, loaded, sched2.get()); EXPECT_EQ(loaded.global_step, 3); - EXPECT_EQ(loaded.consumed_batches, 12); + EXPECT_EQ(loaded.consumed_micro_batches, 12); EXPECT_EQ(sched2->last_step(), sched1->last_step()); EXPECT_NEAR(sched2->learning_rate(), sched1->learning_rate(), kEps); diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index b5352556..8ca00f20 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -20,7 +20,7 @@ class TrainerStateTest : public test::InfiniTrainTest {}; TEST_P(TrainerStateTest, DefaultValues) { TrainerState state; EXPECT_EQ(state.global_step, 0); - EXPECT_EQ(state.consumed_batches, 0); + EXPECT_EQ(state.consumed_micro_batches, 0); EXPECT_EQ(state.n_layer, 0); EXPECT_EQ(state.n_head, 0); EXPECT_EQ(state.n_kv_head, 0); @@ -36,7 +36,7 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { auto dir = std::filesystem::temp_directory_path() / "test_trainer_json"; std::filesystem::remove_all(dir); - TrainerState saved{.global_step = 30, .consumed_batches = 1200}; + TrainerState saved{.global_step = 30, .consumed_micro_batches = 1200}; auto model = std::make_shared(1, 2, true, GetDevice()); auto p = std::make_shared(std::vector{2}, DataType::kFLOAT32, GetDevice()); @@ -51,7 +51,7 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { std::ifstream ifs(dir / "trainer_state.json"); std::string content((std::istreambuf_iterator(ifs)), std::istreambuf_iterator()); EXPECT_NE(content.find("\"global_step\""), std::string::npos); - EXPECT_NE(content.find("\"consumed_batches\""), std::string::npos); + EXPECT_NE(content.find("\"consumed_micro_batches\""), std::string::npos); std::filesystem::remove_all(dir); } @@ -62,7 +62,7 @@ TEST_P(TrainerStateTest, RoundTrip) { TrainerState saved{ .global_step = 99, - .consumed_batches = 5000, + .consumed_micro_batches = 5000, .n_layer = 24, .n_head = 16, .n_kv_head = 8, @@ -92,7 +92,7 @@ TEST_P(TrainerStateTest, RoundTrip) { Checkpoint::Load(dir, *model2, nullptr, loaded, nullptr); EXPECT_EQ(loaded.global_step, 99); - EXPECT_EQ(loaded.consumed_batches, 5000); + EXPECT_EQ(loaded.consumed_micro_batches, 5000); EXPECT_EQ(loaded.n_layer, 24); EXPECT_EQ(loaded.n_head, 16); EXPECT_EQ(loaded.n_kv_head, 8); From f6f47154f98930adf04bc74837d7f5fbea816fc4 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Wed, 5 Aug 2026 11:59:50 +0000 Subject: [PATCH 2/2] feat: resume checkpoints by consumed train samples across batch size changes --- example/gpt2/main.cc | 24 +++++++---------- example/llama3/main.cc | 24 +++++++---------- infini_train/include/checkpoint/checkpoint.h | 2 +- .../include/checkpoint/checkpoint_manager.h | 6 +++-- infini_train/src/checkpoint/checkpoint.cc | 6 ++--- .../src/checkpoint/checkpoint_manager.cc | 22 +++++++++++++--- .../test_checkpoint_serialization.cc | 4 +-- tests/checkpoint/test_lr_scheduler_state.cc | 4 +-- tests/checkpoint/test_trainer_state.cc | 26 +++++++++++++++---- 9 files changed, 69 insertions(+), 49 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 1a77a01b..c06b25ac 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -305,9 +305,9 @@ void Train(const nn::parallel::Rank &rank) { model = std::make_shared(model, rank, ddp_config); } + const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; DistributedDataLoader train_loader(std::make_shared(FLAGS_input_bin, FLAGS_sequence_length), - pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size, - ddp_rank, ddp_world_size); + train_loader_batch_size, ddp_rank, ddp_world_size); std::optional val_loader = std::nullopt; if (!FLAGS_input_val_bin.empty()) { @@ -382,19 +382,13 @@ void Train(const nn::parallel::Rank &rank) { .state = state, .lr_scheduler = scheduler}); start_step = resume_result.global_step; - size_t consumed_micro_batches = resume_result.consumed_micro_batches; + 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_micro_batches > 0) { - const size_t start = train_iter.BatchIndex(); - CHECK(pp_world_size == 1 || consumed_micro_batches % num_micro_batches == 0); - const size_t consumed_loader_batches - = pp_world_size > 1 ? consumed_micro_batches / num_micro_batches : consumed_micro_batches; - const size_t target = consumed_loader_batches + static_cast(ddp_rank); - CHECK_GE(target, start); - CHECK_EQ((target - start) % ddp_world_size, 0); - const size_t num_skips = (target - start) / ddp_world_size; + 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; } } @@ -402,7 +396,7 @@ void Train(const nn::parallel::Rank &rank) { SaveCheckpoint({ .save_dir = save_dir, .global_step = global_step, - .consumed_micro_batches = consumed_micro_batches, + .consumed_train_samples = consumed_train_samples, .n_layer = model_config.n_layer, .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, @@ -474,7 +468,7 @@ void Train(const nn::parallel::Rank &rank) { // 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_micro_batches = train_iter.BatchIndex() - static_cast(ddp_rank); + consumed_train_samples += static_cast(FLAGS_batch_size) * ddp_world_size; x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -508,7 +502,7 @@ void Train(const nn::parallel::Rank &rank) { // 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_micro_batches = (train_iter.BatchIndex() - static_cast(ddp_rank)) * num_micro_batches; + consumed_train_samples += train_loader_batch_size * ddp_world_size; 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 a20a034e..948335dc 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -279,9 +279,9 @@ void Train(const nn::parallel::Rank &rank) { model = std::make_shared(model, rank, ddp_config); } + const size_t train_loader_batch_size = pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size; DistributedDataLoader train_loader(std::make_shared(FLAGS_input_bin, FLAGS_sequence_length), - pp_world_size > 1 ? FLAGS_batch_size * num_micro_batches : FLAGS_batch_size, - ddp_rank, ddp_world_size); + train_loader_batch_size, ddp_rank, ddp_world_size); std::optional val_loader = std::nullopt; if (!FLAGS_input_val_bin.empty()) { @@ -364,19 +364,13 @@ void Train(const nn::parallel::Rank &rank) { .lr_scheduler = scheduler}); start_step = resume_result.global_step; - size_t consumed_micro_batches = resume_result.consumed_micro_batches; + 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_micro_batches > 0) { - const size_t start = train_iter.BatchIndex(); - CHECK(pp_world_size == 1 || consumed_micro_batches % num_micro_batches == 0); - const size_t consumed_loader_batches - = pp_world_size > 1 ? consumed_micro_batches / num_micro_batches : consumed_micro_batches; - const size_t target = consumed_loader_batches + static_cast(ddp_rank); - CHECK_GE(target, start); - CHECK_EQ((target - start) % ddp_world_size, 0); - const size_t num_skips = (target - start) / ddp_world_size; + 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; } } @@ -384,7 +378,7 @@ void Train(const nn::parallel::Rank &rank) { SaveCheckpoint({ .save_dir = save_dir, .global_step = global_step, - .consumed_micro_batches = consumed_micro_batches, + .consumed_train_samples = consumed_train_samples, .n_layer = model_config.n_layer, .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, @@ -454,7 +448,7 @@ void Train(const nn::parallel::Rank &rank) { // 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_micro_batches = train_iter.BatchIndex() - static_cast(ddp_rank); + consumed_train_samples += static_cast(FLAGS_batch_size) * ddp_world_size; x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); @@ -487,7 +481,7 @@ void Train(const nn::parallel::Rank &rank) { // 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_micro_batches = (train_iter.BatchIndex() - static_cast(ddp_rank)) * num_micro_batches; + consumed_train_samples += train_loader_batch_size * ddp_world_size; x = std::make_shared(x->To(device)); y = std::make_shared(y->To(device)); diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index ecb5ce16..a69aad23 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -17,7 +17,7 @@ class Module; struct TrainerState { int64_t global_step = 0; - int64_t consumed_micro_batches = 0; + int64_t consumed_train_samples = 0; int64_t n_layer = 0; int64_t n_head = 0; int64_t n_kv_head = 0; diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 92a54632..13490ccc 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -34,13 +34,13 @@ struct ResumeFromCheckpointArgs { struct ResumeFromCheckpointResult { int global_step = 0; - size_t consumed_micro_batches = 0; + size_t consumed_train_samples = 0; }; struct SaveCheckpointArgs { std::filesystem::path save_dir; int64_t global_step = 0; - size_t consumed_micro_batches = 0; + size_t consumed_train_samples = 0; int64_t n_layer = 0; int64_t n_head = 0; int64_t n_kv_head = 0; @@ -61,3 +61,5 @@ struct SaveCheckpointArgs { ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args); void SaveCheckpoint(const SaveCheckpointArgs &args); + +size_t DataLoaderBatchesToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size); diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index 38cb5c19..ad51bab1 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -237,7 +237,7 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m } LOG(ERROR) << "[CKPT] Load done: global_step=" << state.global_step - << ", consumed_micro_batches =" << state.consumed_micro_batches << ", topology(ddp,tp,sp,pp)=(" + << ", consumed_train_samples=" << state.consumed_train_samples << ", topology(ddp,tp,sp,pp)=(" << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; } @@ -320,7 +320,7 @@ void Checkpoint::SaveTrainerState(const std::filesystem::path &path, const Train ofs << " \"n_embd\": " << state.n_embd << ",\n"; ofs << " \"vocab_size\": " << state.vocab_size << ",\n"; ofs << " \"global_step\": " << state.global_step << ",\n"; - ofs << " \"consumed_micro_batches\": " << state.consumed_micro_batches << ",\n"; + ofs << " \"consumed_train_samples\": " << state.consumed_train_samples << ",\n"; ofs << " \"ddp_size\": " << state.ddp_size << ",\n"; ofs << " \"tp_size\": " << state.tp_size << ",\n"; ofs << " \"sp_size\": " << state.sp_size << ",\n"; @@ -341,7 +341,7 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.n_embd = ExtractNumberField(content, "n_embd", 0); state.vocab_size = ExtractNumberField(content, "vocab_size", 0); state.global_step = ExtractNumberField(content, "global_step", 0); - state.consumed_micro_batches = ExtractNumberField(content, "consumed_micro_batches", 0); + state.consumed_train_samples = ExtractNumberField(content, "consumed_train_samples", 0); state.ddp_size = ExtractNumberField(content, "ddp_size", 1); state.tp_size = ExtractNumberField(content, "tp_size", 1); state.sp_size = ExtractNumberField(content, "sp_size", 1); diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index 427ec2f4..c6e31cdd 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -4,6 +4,7 @@ #include #include #include +#include #include #include #include @@ -63,10 +64,10 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & CHECK_EQ(args.state.pp_size, pp_world_size) << "PP size mismatch: checkpoint has PP=" << args.state.pp_size << ", but current run has PP=" << pp_world_size; - result.consumed_micro_batches = static_cast(std::max(args.state.consumed_micro_batches, 0)); + result.consumed_train_samples = static_cast(std::max(args.state.consumed_train_samples, 0)); if (args.rank.IsMainRank()) { - LOG(INFO) << std::format("Resume training from step {}, consumed_micro_batches {}", args.state.global_step, - args.state.consumed_micro_batches); + LOG(INFO) << std::format("Resume training from step {}, consumed_train_samples {}", args.state.global_step, + args.state.consumed_train_samples); } return result; @@ -77,7 +78,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { TrainerState state; state.global_step = args.global_step; - state.consumed_micro_batches = static_cast(args.consumed_micro_batches); + state.consumed_train_samples = static_cast(args.consumed_train_samples); state.n_layer = args.n_layer; state.n_head = args.n_head; state.n_kv_head = args.n_kv_head; @@ -118,3 +119,16 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { } } } + +size_t DataLoaderBatchesToSkip(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) + << "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; +} diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index c03b8462..c44595d1 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -29,7 +29,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { *model1->mutable_parameter("bias") = p2; auto opt1 = std::make_shared(model1->Parameters(), 0.01); - TrainerState saved{.global_step = 42, .consumed_micro_batches = 100}; + TrainerState saved{.global_step = 42, .consumed_train_samples = 100}; Checkpoint::Save(dir, *model1, opt1.get(), saved, nullptr); auto model2 = std::make_shared(3, 2, true, GetDevice()); @@ -45,7 +45,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { Checkpoint::Load(dir, *model2, opt2.get(), loaded, nullptr); EXPECT_EQ(loaded.global_step, 42); - EXPECT_EQ(loaded.consumed_micro_batches, 100); + EXPECT_EQ(loaded.consumed_train_samples, 100); auto w1_cpu = model2->parameter("weight")->To(Device()); const float *data = static_cast(w1_cpu.DataPtr()); diff --git a/tests/checkpoint/test_lr_scheduler_state.cc b/tests/checkpoint/test_lr_scheduler_state.cc index 9df7238b..a0e07dd9 100644 --- a/tests/checkpoint/test_lr_scheduler_state.cc +++ b/tests/checkpoint/test_lr_scheduler_state.cc @@ -59,7 +59,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { auto sched1 = CreateLRScheduler(opt1, MakeSchedulerConfig()); StepTimes(sched1, 3); - TrainerState saved{.global_step = 3, .consumed_micro_batches = 12}; + TrainerState saved{.global_step = 3, .consumed_train_samples = 12}; Checkpoint::Save(dir, *model1, nullptr, saved, sched1.get()); EXPECT_TRUE(std::filesystem::exists(dir / "lr_scheduler.ckpt")); @@ -71,7 +71,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { Checkpoint::Load(dir, *model2, nullptr, loaded, sched2.get()); EXPECT_EQ(loaded.global_step, 3); - EXPECT_EQ(loaded.consumed_micro_batches, 12); + EXPECT_EQ(loaded.consumed_train_samples, 12); EXPECT_EQ(sched2->last_step(), sched1->last_step()); EXPECT_NEAR(sched2->learning_rate(), sched1->learning_rate(), kEps); diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index 8ca00f20..ec4d61e8 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -5,6 +5,7 @@ #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/checkpoint_manager.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/optimizer.h" @@ -20,7 +21,7 @@ class TrainerStateTest : public test::InfiniTrainTest {}; TEST_P(TrainerStateTest, DefaultValues) { TrainerState state; EXPECT_EQ(state.global_step, 0); - EXPECT_EQ(state.consumed_micro_batches, 0); + EXPECT_EQ(state.consumed_train_samples, 0); EXPECT_EQ(state.n_layer, 0); EXPECT_EQ(state.n_head, 0); EXPECT_EQ(state.n_kv_head, 0); @@ -36,7 +37,7 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { auto dir = std::filesystem::temp_directory_path() / "test_trainer_json"; std::filesystem::remove_all(dir); - TrainerState saved{.global_step = 30, .consumed_micro_batches = 1200}; + TrainerState saved{.global_step = 30, .consumed_train_samples = 1200}; auto model = std::make_shared(1, 2, true, GetDevice()); auto p = std::make_shared(std::vector{2}, DataType::kFLOAT32, GetDevice()); @@ -51,7 +52,7 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { std::ifstream ifs(dir / "trainer_state.json"); std::string content((std::istreambuf_iterator(ifs)), std::istreambuf_iterator()); EXPECT_NE(content.find("\"global_step\""), std::string::npos); - EXPECT_NE(content.find("\"consumed_micro_batches\""), std::string::npos); + EXPECT_NE(content.find("\"consumed_train_samples\""), std::string::npos); std::filesystem::remove_all(dir); } @@ -62,7 +63,7 @@ TEST_P(TrainerStateTest, RoundTrip) { TrainerState saved{ .global_step = 99, - .consumed_micro_batches = 5000, + .consumed_train_samples = 5000, .n_layer = 24, .n_head = 16, .n_kv_head = 8, @@ -92,7 +93,7 @@ TEST_P(TrainerStateTest, RoundTrip) { Checkpoint::Load(dir, *model2, nullptr, loaded, nullptr); EXPECT_EQ(loaded.global_step, 99); - EXPECT_EQ(loaded.consumed_micro_batches, 5000); + EXPECT_EQ(loaded.consumed_train_samples, 5000); EXPECT_EQ(loaded.n_layer, 24); EXPECT_EQ(loaded.n_head, 16); EXPECT_EQ(loaded.n_kv_head, 8); @@ -104,4 +105,19 @@ 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), + 50); + EXPECT_EQ(DataLoaderBatchesToSkip(/*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), + "does not align"); +} + INFINI_TRAIN_REGISTER_TEST(TrainerStateTest);