From 9626bdb5dae4fa48245d9b8df9c4c1de403d25ad Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Wed, 12 Aug 2026 07:29:14 +0000 Subject: [PATCH 1/5] feat: reshard distributed checkpoints across TP and PP --- infini_train/include/checkpoint/checkpoint.h | 69 ++- .../include/checkpoint/checkpoint_manager.h | 1 - .../include/checkpoint/load_planner.h | 52 ++ .../include/checkpoint/load_strategy.h | 30 ++ infini_train/include/checkpoint/reshard.h | 22 + .../include/checkpoint/save_planner.h | 79 +++ infini_train/include/checkpoint/shard_spec.h | 56 +++ infini_train/include/nn/modules/module.h | 10 +- .../transformer/causal_self_attention.h | 2 + .../nn/modules/transformer/transformer.h | 5 + .../parallel/ddp/distributed_data_parallel.h | 5 + .../nn/parallel/pp/pipeline_parallel.h | 6 + .../include/nn/parallel/tensor_parallel.h | 7 + infini_train/src/checkpoint/checkpoint.cc | 473 +++++++++++++++++- .../src/checkpoint/checkpoint_manager.cc | 261 +++++++--- infini_train/src/checkpoint/load_planner.cc | 244 +++++++++ infini_train/src/checkpoint/load_strategy.cc | 129 +++++ infini_train/src/checkpoint/reshard.cc | 71 +++ infini_train/src/checkpoint/save_planner.cc | 75 +++ infini_train/src/nn/modules/module.cc | 38 ++ .../transformer/causal_self_attention.cc | 34 ++ .../src/nn/modules/transformer/transformer.cc | 77 +++ .../parallel/ddp/distributed_data_parallel.cc | 18 + .../src/nn/parallel/pp/pipeline_parallel.cc | 17 + .../src/nn/parallel/tensor_parallel.cc | 81 ++- .../test_checkpoint_serialization.cc | 400 +++++++++++++++ tests/checkpoint/test_optimizer_state.cc | 19 + .../test_transformer_architecture.cc | 4 + 28 files changed, 2198 insertions(+), 87 deletions(-) create mode 100644 infini_train/include/checkpoint/load_planner.h create mode 100644 infini_train/include/checkpoint/load_strategy.h create mode 100644 infini_train/include/checkpoint/reshard.h create mode 100644 infini_train/include/checkpoint/save_planner.h create mode 100644 infini_train/include/checkpoint/shard_spec.h create mode 100644 infini_train/src/checkpoint/load_planner.cc create mode 100644 infini_train/src/checkpoint/load_strategy.cc create mode 100644 infini_train/src/checkpoint/reshard.cc create mode 100644 infini_train/src/checkpoint/save_planner.cc diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index a69aad232..586c1e26c 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -6,6 +6,11 @@ #include #include #include +#include + +#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/lr_scheduler.h" namespace infini_train { class Optimizer; @@ -37,9 +42,69 @@ class Checkpoint { static void Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, LRScheduler *lr_scheduler); + static void SaveSharded(const std::filesystem::path &checkpoint_dir, const checkpoint::ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + const TrainerState &state, int global_rank); + + static void SaveStateDictFile(const std::filesystem::path &path, + const std::unordered_map> &state_dict); + + static std::unordered_map> + LoadStateDictFile(const std::filesystem::path &path); + + struct CheckpointMetadata { + int version = 0; + int64_t iteration = 0; + + struct ParallelConfig { + int tp_size = 1; + int pp_size = 1; + int dp_size = 1; + int sp_size = 1; + } parallel_config; + + struct TensorEntry { + std::string key; + std::string dtype_str; + std::vector global_shape; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + std::vector segments; + std::string file; + uint64_t offset = 0; + uint64_t byte_size = 0; + std::vector stored_on_ranks; + int pp_rank = 0; + }; + + std::vector tensors; + bool has_metadata = false; + }; + + static CheckpointMetadata LoadMetadata(const std::filesystem::path &checkpoint_dir); + static void SaveMetadataFile(const std::filesystem::path &path, const CheckpointMetadata &metadata); + + // Public LR-scheduler serialization helpers used by checkpoint_manager. + static void SaveLRSchedulerStateFile(const std::filesystem::path &path, const LRSchedulerStateDict &state_dict); + static LRSchedulerStateDict LoadLRSchedulerStateFile(const std::filesystem::path &path); + + // Public trainer-state serialization helpers used by checkpoint_manager. + static void SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state); + static TrainerState LoadTrainerStateFile(const std::filesystem::path &path); + private: - static void SaveStateDict(const std::filesystem::path &path, - const std::unordered_map> &state_dict); + struct SavedTensorLocation { + uint64_t data_offset = 0; + uint64_t byte_size = 0; + }; + using SavedTensorLocations = std::unordered_map; + + static SavedTensorLocations + SaveStateDict(const std::filesystem::path &path, + const std::unordered_map> &state_dict); static std::unordered_map> LoadStateDict(const std::filesystem::path &path); diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 13490ccce..47e07fdd6 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -6,7 +6,6 @@ #include #include "infini_train/include/checkpoint/checkpoint.h" -#include "infini_train/include/dataloader.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/rank.h" #include "infini_train/include/optimizer.h" diff --git a/infini_train/include/checkpoint/load_planner.h b/infini_train/include/checkpoint/load_planner.h new file mode 100644 index 000000000..7011cc730 --- /dev/null +++ b/infini_train/include/checkpoint/load_planner.h @@ -0,0 +1,52 @@ +#pragma once + +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/datatype.h" + +namespace infini_train::checkpoint { + +// One storage-region transfer from a saved shard into a target local tensor. +struct ReadItem { + std::string key; + std::string filename; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + uint64_t byte_size = 0; + uint64_t data_offset = 0; + int shard_dim = -1; + int64_t source_offset = 0; + int64_t target_offset = 0; + int64_t length = 0; + std::vector source_shape; +}; + +// All reads required to materialize one target local tensor. +struct TargetTensorPlan { + std::string key; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector target_shape; + int shard_dim = -1; + int64_t trailing_zero_fill = 0; + std::vector reads; +}; + +// Complete load plan for one rank. +struct LoadPlan { + std::map tensors; +}; + +class LoadPlanner { +public: + // Compute saved-to-target overlaps from explicit global shard coordinates. + static LoadPlan PlanReshard(const Checkpoint::CheckpointMetadata &metadata, + const ShardedStateDict &target_state_dict); +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/load_strategy.h b/infini_train/include/checkpoint/load_strategy.h new file mode 100644 index 000000000..f5f6e2c0f --- /dev/null +++ b/infini_train/include/checkpoint/load_strategy.h @@ -0,0 +1,30 @@ +#pragma once + +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/load_planner.h" + +namespace infini_train { +class Tensor; +} + +namespace infini_train::checkpoint { + +using LoadedStateDict = std::unordered_map>; + +class LoadStrategy { +public: + virtual ~LoadStrategy() = default; + virtual LoadedStateDict Execute(const std::filesystem::path &checkpoint_dir, const LoadPlan &plan) = 0; +}; + +/// Reads source regions directly from metadata offsets while caching one open stream per file. +class IndexedRegionLoadStrategy final : public LoadStrategy { +public: + LoadedStateDict Execute(const std::filesystem::path &checkpoint_dir, const LoadPlan &plan) override; +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/reshard.h b/infini_train/include/checkpoint/reshard.h new file mode 100644 index 000000000..c091d2757 --- /dev/null +++ b/infini_train/include/checkpoint/reshard.h @@ -0,0 +1,22 @@ +#pragma once + +#include + +#include "infini_train/include/checkpoint/checkpoint.h" + +namespace infini_train { +class LRScheduler; +class Optimizer; +namespace nn { +class Module; +} +} // namespace infini_train + +namespace infini_train::checkpoint { + +// Restore this rank's target model shards from a distributed checkpoint. +void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, LRScheduler *lr_scheduler, + const Checkpoint::CheckpointMetadata &metadata); + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/save_planner.h b/infini_train/include/checkpoint/save_planner.h new file mode 100644 index 000000000..02a5a21c1 --- /dev/null +++ b/infini_train/include/checkpoint/save_planner.h @@ -0,0 +1,79 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/datatype.h" + +namespace infini_train { +class Tensor; +} + +namespace infini_train::checkpoint { + +// Physical write description for one local tensor shard. +struct WriteItem { + std::string key; + std::string filename; // "model.ckpt" or "optimizer.ckpt" + uint64_t offset = 0; // Planned byte offset in the checkpoint file. + uint64_t byte_size = 0; // Tensor payload size in bytes. + DataType dtype = DataType::kFLOAT32; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + int rank = 0; +}; + +// Build the local tensor write layout from a ShardedStateDict. +class SavePlanner { +public: + static std::vector Plan(const ShardedStateDict &sd, int rank); +}; + +ShardedStateDict +BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, + const std::unordered_map> &optimizer_state); + +// Return the number of payload bytes required by a tensor. +inline uint64_t TensorByteSize(DataType dtype, const std::vector &shape) { + uint64_t numel = 1; + for (auto d : shape) { numel *= static_cast(d); } + switch (dtype) { + case DataType::kBFLOAT16: + case DataType::kFLOAT16: + return numel * 2; + case DataType::kFLOAT32: + return numel * 4; + case DataType::kFLOAT64: + case DataType::kINT64: + case DataType::kUINT64: + return numel * 8; + case DataType::kINT32: + case DataType::kUINT32: + return numel * 4; + case DataType::kINT16: + case DataType::kUINT16: + return numel * 2; + case DataType::kINT8: + case DataType::kUINT8: + case DataType::kBOOL: + return numel; + default: + return numel * 4; + } +} + +// Compute one rank's balanced interval, including non-divisible dimensions. +inline std::pair GetRankSliceRange(int64_t global_size, int world_size, int rank) { + int64_t per_rank = global_size / world_size; + int64_t remainder = global_size % world_size; + int64_t start = rank * per_rank + std::min(rank, remainder); + int64_t local_size = per_rank + (rank < remainder ? 1 : 0); + return {start, local_size}; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/shard_spec.h b/infini_train/include/checkpoint/shard_spec.h new file mode 100644 index 000000000..75263765a --- /dev/null +++ b/infini_train/include/checkpoint/shard_spec.h @@ -0,0 +1,56 @@ +#pragma once + +#include +#include +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/datatype.h" + +namespace infini_train::checkpoint { + +struct ShardSegment { + int64_t global_offset = 0; + int64_t local_offset = 0; + int64_t length = 0; + + bool operator==(const ShardSegment &other) const = default; +}; + +// Logical tensor shard metadata, aligned with Megatron-LM's ShardedTensor model. +struct ShardedTensor { + std::string key; + std::string local_key; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + // Optional disjoint regions along the single fragmented axis. This is used + // by layouts such as rank-local [Q, K, V], which are not one contiguous + // slice of the logical global [Q, K, V] tensor. + std::vector segments; + + bool operator==(const ShardedTensor &other) const { + return key == other.key && local_key == other.local_key && dtype == other.dtype + && global_shape == other.global_shape && local_shape == other.local_shape + && global_offset == other.global_offset && axis_fragmentations == other.axis_fragmentations + && segments == other.segments; + } +}; + +struct ShardedStateDict { + std::map tensors; + + void Merge(ShardedStateDict &&other) { + for (auto &[key, info] : other.tensors) { + const auto display_key = key; + const auto [_, inserted] = tensors.emplace(std::move(key), std::move(info)); + CHECK(inserted) << "Duplicate sharded state-dict key: " << display_key; + } + } +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 8570c4768..0816d247f 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -6,6 +6,7 @@ #include #include +#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/datatype.h" #include "infini_train/include/device.h" @@ -63,7 +64,7 @@ class Module : public std::enable_shared_from_this { // InfiniTrain's NamedParameters returns results ordered by full parameter name. // TODO: Align with PyTorch's ordering in the future. - std::vector>> + virtual std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const; bool has_parameter(const std::string &name) const; std::shared_ptr *mutable_parameter(const std::string &name); @@ -75,11 +76,14 @@ class Module : public std::enable_shared_from_this { std::shared_ptr &mutable_module(const std::string &name); const Module &module(const std::string &name) const; - std::unordered_map> StateDict() const; + virtual std::unordered_map> StateDict() const; + + // Return state-dict metadata with global shard coordinates. + virtual checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const; // Current behavior: missing keys / shape / dtype mismatches are FATAL errors; unexpected keys in state_dict are // WARNING-only and silently ignored. - void LoadStateDict(const std::unordered_map> &state_dict); + virtual void LoadStateDict(const std::unordered_map> &state_dict); // operator() calls hooks and Forward std::vector> operator()(const std::vector> &input_tensors); diff --git a/infini_train/include/nn/modules/transformer/causal_self_attention.h b/infini_train/include/nn/modules/transformer/causal_self_attention.h index a3fa6d257..6e713b19f 100644 --- a/infini_train/include/nn/modules/transformer/causal_self_attention.h +++ b/infini_train/include/nn/modules/transformer/causal_self_attention.h @@ -25,6 +25,8 @@ class CausalSelfAttention : public infini_train::nn::CloneableModule> Forward(const std::vector> &x) override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + private: TransformerConfig config_; int64_t n_head_ = 0; diff --git a/infini_train/include/nn/modules/transformer/transformer.h b/infini_train/include/nn/modules/transformer/transformer.h index 0471c32fe..455f37c28 100644 --- a/infini_train/include/nn/modules/transformer/transformer.h +++ b/infini_train/include/nn/modules/transformer/transformer.h @@ -78,6 +78,11 @@ class TransformerModel : public CloneableModule { const TransformerConfig &Config() const { return config_; } + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; + void LoadStateDict(const std::unordered_map> &state_dict) override; + private: const TransformerConfig config_; const infini_train::nn::parallel::StageInfo stage_info_; diff --git a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h index 905816b56..4aa130dc8 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -31,6 +31,11 @@ class DistributedDataParallel : public nn::Module { std::vector> Forward(const std::vector> &input_tensors) override; std::shared_ptr module() const; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; + std::unordered_map> StateDict() const override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + void LoadStateDict(const std::unordered_map> &state_dict) override; std::unique_ptr no_sync() override; diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 25939bdc2..58f48cd59 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_parallel.h +++ b/infini_train/include/nn/parallel/pp/pipeline_parallel.h @@ -40,6 +40,12 @@ class PipelineParallel : public Module { std::vector> *mutable_chunks(); + std::unordered_map> StateDict() const override; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + void LoadStateDict(const std::unordered_map> &state_dict) override; + private: void BuildPipelineStage(const std::vector> &recv_shape, Device device, std::vector> &&chunks); diff --git a/infini_train/include/nn/parallel/tensor_parallel.h b/infini_train/include/nn/parallel/tensor_parallel.h index a44aceeaf..06a0abea6 100644 --- a/infini_train/include/nn/parallel/tensor_parallel.h +++ b/infini_train/include/nn/parallel/tensor_parallel.h @@ -4,6 +4,7 @@ #include #include "infini_train/include/autograd/function.h" +#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/process_group.h" @@ -37,6 +38,8 @@ class ColumnParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + protected: bool bias_ = true; bool gather_output_ = false; // whether to return full local output tensor after forward (need gather) @@ -66,6 +69,8 @@ class RowParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + protected: bool bias_ = true; bool reduce_output_ = false; // whether to return full local output tensor after forward (need reduce) @@ -85,6 +90,8 @@ class VocabParallelEmbedding : public nn::CloneableModule> Forward(const std::vector> &input_tensors) override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + private: bool reduce_scatter_embeddings_ = false; // whether to perform ReduceScatter after embedding lookup diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index ad51bab1c..ef3559fb4 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -1,5 +1,6 @@ #include "infini_train/include/checkpoint/checkpoint.h" +#include #include #include #include @@ -11,8 +12,10 @@ #include "glog/logging.h" +#include "infini_train/include/checkpoint/save_planner.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -241,13 +244,15 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; } -void Checkpoint::SaveStateDict(const std::filesystem::path &path, - const std::unordered_map> &state_dict) { +Checkpoint::SavedTensorLocations +Checkpoint::SaveStateDict(const std::filesystem::path &path, + const std::unordered_map> &state_dict) { std::ofstream ofs(path, std::ios::binary); CHECK(ofs.is_open()) << "Failed to open checkpoint file: " << path; uint32_t magic = kCkptMagic; uint32_t version = kCkptVersion; + SavedTensorLocations locations; uint32_t count = static_cast(state_dict.size()); ofs.write(reinterpret_cast(&magic), sizeof(magic)); ofs.write(reinterpret_cast(&version), sizeof(version)); @@ -267,8 +272,14 @@ void Checkpoint::SaveStateDict(const std::filesystem::path &path, Tensor cpu_tensor = tensor->To(Device()); uint64_t bytes = static_cast(cpu_tensor.SizeInBytes()); ofs.write(reinterpret_cast(&bytes), sizeof(bytes)); + const auto data_offset = ofs.tellp(); + CHECK(data_offset != std::streampos(-1)) << "Failed to record tensor offset for " << name; + locations.emplace( + name, SavedTensorLocation{.data_offset = static_cast(static_cast(data_offset)), + .byte_size = bytes}); ofs.write(reinterpret_cast(cpu_tensor.DataPtr()), static_cast(bytes)); } + return locations; } std::unordered_map> Checkpoint::LoadStateDict(const std::filesystem::path &path) { @@ -348,4 +359,462 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.pp_size = ExtractNumberField(content, "pp_size", 1); return state; } + +void Checkpoint::SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state) { + SaveTrainerState(path, state); +} + +TrainerState Checkpoint::LoadTrainerStateFile(const std::filesystem::path &path) { return LoadTrainerState(path); } + +void Checkpoint::SaveLRSchedulerStateFile(const std::filesystem::path &path, const LRSchedulerStateDict &state_dict) { + SaveLRSchedulerState(path, state_dict); +} + +LRSchedulerStateDict Checkpoint::LoadLRSchedulerStateFile(const std::filesystem::path &path) { + return LoadLRSchedulerState(path); +} + +void Checkpoint::SaveStateDictFile(const std::filesystem::path &path, + const std::unordered_map> &state_dict) { + SaveStateDict(path, state_dict); +} + +std::unordered_map> +Checkpoint::LoadStateDictFile(const std::filesystem::path &path) { + return LoadStateDict(path); +} + +// ----------------------------------------------------------------------------- +// Save local shards and a temporary rank manifest from a ShardedStateDict. +// ----------------------------------------------------------------------------- + +static std::string DataTypeToString(DataType dt) { + auto it = kDataTypeToDesc.find(dt); + if (it != kDataTypeToDesc.end()) { + return it->second; + } + return "fp32"; +} + +void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, + const checkpoint::ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + const TrainerState &state, int global_rank) { + std::filesystem::create_directories(checkpoint_dir); + LOG(INFO) << "[CKPT] SaveSharded begin: dir=" << checkpoint_dir << ", global_step=" << state.global_step + << ", rank=" << global_rank; + + SavedTensorLocations model_file_index; + SavedTensorLocations optimizer_file_index; + + // Save model tensors separately from optimizer tensors. + { + std::unordered_map> filtered_sd; + for (const auto &[key, info] : sharded_sd.tensors) { + // Optimizer tensors are serialized separately. + if (key.starts_with("adam.")) { + continue; + } + // Match metadata keys to the local tensor payloads. + const auto &local_key = info.local_key.empty() ? key : info.local_key; + auto it = state_dict.find(local_key); + if (it != state_dict.end()) { + filtered_sd.emplace(key, it->second); + } + } + if (!filtered_sd.empty()) { + model_file_index = SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); + } + } + + // Save the rank-local optimizer state. + if (!optimizer_state.empty()) { + optimizer_file_index = SaveStateDict(checkpoint_dir / "optimizer.ckpt", optimizer_state); + } + + // Write the temporary rank manifest. + { + std::ofstream ofs(checkpoint_dir / "metadata.json"); + CHECK(ofs.is_open()) << "Failed to open metadata.json: " << checkpoint_dir / "metadata.json"; + + ofs << "{\n"; + ofs << " \"version\": 3,\n"; + ofs << " \"format\": \"infinitrain_sharded\",\n"; + ofs << " \"iteration\": " << state.global_step << ",\n"; + ofs << " \"parallel_config\": {\n"; + ofs << " \"tp_size\": " << state.tp_size << ",\n"; + ofs << " \"pp_size\": " << state.pp_size << ",\n"; + ofs << " \"dp_size\": " << state.ddp_size << ",\n"; + ofs << " \"sp_size\": " << state.sp_size << "\n"; + ofs << " },\n"; + ofs << " \"model_config\": {\n"; + ofs << " \"n_layer\": " << state.n_layer << ",\n"; + ofs << " \"n_head\": " << state.n_head << ",\n"; + ofs << " \"n_kv_head\": " << state.n_kv_head << ",\n"; + ofs << " \"n_embd\": " << state.n_embd << ",\n"; + ofs << " \"vocab_size\": " << state.vocab_size << "\n"; + ofs << " },\n"; + ofs << " \"tensors\": [\n"; + + std::vector emitted_items; + for (const auto &item : write_items) { + if (sharded_sd.tensors.contains(item.key)) { + emitted_items.push_back(&item); + } + } + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(global_rank, dp_rank, tp_rank, pp_rank); + for (size_t i = 0; i < emitted_items.size(); ++i) { + const auto &item = *emitted_items[i]; + const auto it = sharded_sd.tensors.find(item.key); + const auto &file_index = item.filename == "optimizer.ckpt" ? optimizer_file_index : model_file_index; + const auto storage_it = file_index.find(item.key); + CHECK(storage_it != file_index.end()) << "Missing stored tensor metadata for " << item.key; + const auto &storage = storage_it->second; + CHECK_EQ(storage.byte_size, item.byte_size); + + ofs << " {\n"; + ofs << " \"key\": \"" << item.key << "\",\n"; + ofs << " \"dtype\": \"" << DataTypeToString(item.dtype) << "\",\n"; + + // global_shape + ofs << " \"global_shape\": ["; + const auto &gs = it->second.global_shape; + for (size_t j = 0; j < gs.size(); ++j) { ofs << gs[j] << (j + 1 < gs.size() ? ", " : ""); } + ofs << "],\n"; + + ofs << " \"local_shape\": ["; + const auto &ls = it->second.local_shape; + for (size_t j = 0; j < ls.size(); ++j) { ofs << ls[j] << (j + 1 < ls.size() ? ", " : ""); } + ofs << "],\n"; + + ofs << " \"global_offset\": ["; + for (size_t j = 0; j < it->second.global_offset.size(); ++j) { + ofs << it->second.global_offset[j] << (j + 1 < it->second.global_offset.size() ? ", " : ""); + } + ofs << "],\n"; + ofs << " \"axis_fragmentations\": ["; + for (size_t j = 0; j < it->second.axis_fragmentations.size(); ++j) { + ofs << it->second.axis_fragmentations[j] << (j + 1 < it->second.axis_fragmentations.size() ? ", " : ""); + } + ofs << "],\n"; + auto write_segments = [&](const char *name, auto member) { + ofs << " \"" << name << "\": ["; + for (size_t j = 0; j < it->second.segments.size(); ++j) { + ofs << it->second.segments[j].*member << (j + 1 < it->second.segments.size() ? ", " : ""); + } + ofs << "],\n"; + }; + write_segments("segment_global_offsets", &checkpoint::ShardSegment::global_offset); + write_segments("segment_local_offsets", &checkpoint::ShardSegment::local_offset); + write_segments("segment_lengths", &checkpoint::ShardSegment::length); + + ofs << " \"file\": \"" << item.filename << "\",\n"; + ofs << " \"offset\": " << storage.data_offset << ",\n"; + ofs << " \"byte_size\": " << item.byte_size << ",\n"; + ofs << " \"pp_rank\": " << pp_rank << ",\n"; + ofs << " \"stored_on_ranks\": [" << global_rank << "]\n"; + ofs << " }"; + if (i + 1 < emitted_items.size()) { + ofs << ","; + } + ofs << "\n"; + } + + ofs << " ]\n"; + ofs << "}\n"; + + LOG(INFO) << "[CKPT] metadata.json written"; + } + + LOG(ERROR) << "[CKPT] SaveSharded done: dir=" << checkpoint_dir; +} + +// Load one manifest or aggregate writer manifests while finalizing a checkpoint. +static std::string ExtractJsonString(const std::string &obj, const std::string &key) { + auto token = std::string("\"") + key + "\""; + auto pos = obj.find(token); + if (pos == std::string::npos) { + return ""; + } + auto q1 = obj.find('"', pos + token.size()); + if (q1 == std::string::npos) { + return ""; + } + auto q2 = obj.find('"', q1 + 1); + if (q2 == std::string::npos) { + return ""; + } + return obj.substr(q1 + 1, q2 - q1 - 1); +} + +static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem::path &checkpoint_dir) { + Checkpoint::CheckpointMetadata meta; + auto metadata_path = checkpoint_dir / "metadata.json"; + if (!std::filesystem::exists(metadata_path)) { + meta.has_metadata = false; + return meta; + } + + std::ifstream ifs(metadata_path); + CHECK(ifs.is_open()) << "Failed to open metadata.json: " << metadata_path; + const std::string content((std::istreambuf_iterator(ifs)), std::istreambuf_iterator()); + + meta.has_metadata = true; + meta.version = ExtractNumberField(content, "version", 0); + meta.iteration = ExtractNumberField(content, "iteration", 0); + meta.parallel_config.tp_size = ExtractNumberField(content, "tp_size", 1); + meta.parallel_config.pp_size = ExtractNumberField(content, "pp_size", 1); + meta.parallel_config.dp_size = ExtractNumberField(content, "dp_size", 1); + meta.parallel_config.sp_size = ExtractNumberField(content, "sp_size", 1); + + // Locate the tensors array. + auto tensors_key = content.find("\"tensors\""); + if (tensors_key == std::string::npos) { + return meta; + } + + auto array_start = content.find('[', tensors_key); + if (array_start == std::string::npos) { + return meta; + } + + int depth = 1; + size_t pos = array_start + 1; + while (pos < content.size() && depth > 0) { + if (content[pos] == '[') { + ++depth; + } else if (content[pos] == ']') { + --depth; + } + ++pos; + } + std::string tensor_block = content.substr(array_start + 1, pos - array_start - 2); + + // Parse each tensor object. + size_t obj_pos = 0; + while ((obj_pos = tensor_block.find('{', obj_pos)) != std::string::npos) { + int object_depth = 1; + size_t obj_end = obj_pos + 1; + while (obj_end < tensor_block.size() && object_depth > 0) { + if (tensor_block[obj_end] == '{') { + ++object_depth; + } + if (tensor_block[obj_end] == '}') { + --object_depth; + } + ++obj_end; + } + if (obj_end > 0) { + --obj_end; + } + if (obj_end == std::string::npos) { + break; + } + + std::string obj = tensor_block.substr(obj_pos, obj_end - obj_pos + 1); + + Checkpoint::CheckpointMetadata::TensorEntry entry; + entry.key = ExtractJsonString(obj, "key"); + entry.file = ExtractJsonString(obj, "file"); + entry.dtype_str = ExtractJsonString(obj, "dtype"); + entry.offset = ExtractNumberField(obj, "offset", 0); + entry.byte_size = ExtractNumberField(obj, "byte_size", 0); + entry.pp_rank = ExtractNumberField(obj, "pp_rank", 0); + + // global_shape: [x, y, z] + auto gs_pos = obj.find("\"global_shape\""); + if (gs_pos != std::string::npos) { + auto b1 = obj.find('[', gs_pos); + auto b2 = obj.find(']', b1); + if (b1 != std::string::npos && b2 != std::string::npos) { + std::string gs = obj.substr(b1 + 1, b2 - b1 - 1); + std::stringstream ss(gs); + std::string tok; + while (std::getline(ss, tok, ',')) { + try { + entry.global_shape.push_back(std::stoll(tok)); + } catch (...) {} + } + } + } + + auto ls_pos = obj.find("\"local_shape\""); + if (ls_pos != std::string::npos) { + auto b1 = obj.find('[', ls_pos); + auto b2 = obj.find(']', b1); + std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); + std::string tok; + while (std::getline(ss, tok, ',')) { + try { + entry.local_shape.push_back(std::stoll(tok)); + } catch (...) {} + } + } + + auto offset_pos = obj.find("\"global_offset\""); + if (offset_pos != std::string::npos) { + auto b1 = obj.find('[', offset_pos); + auto b2 = obj.find(']', b1); + std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); + std::string token; + while (std::getline(ss, token, ',')) { + try { + entry.global_offset.push_back(std::stoll(token)); + } catch (...) {} + } + } + + auto fragments_pos = obj.find("\"axis_fragmentations\""); + if (fragments_pos != std::string::npos) { + auto b1 = obj.find('[', fragments_pos); + auto b2 = obj.find(']', b1); + std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); + std::string token; + while (std::getline(ss, token, ',')) { + try { + entry.axis_fragmentations.push_back(std::stoi(token)); + } catch (...) {} + } + } + + auto extract_int64_array = [&](const char *name) { + std::vector values; + auto field_pos = obj.find(std::string("\"") + name + "\""); + if (field_pos == std::string::npos) { + return values; + } + auto b1 = obj.find('[', field_pos); + auto b2 = obj.find(']', b1); + if (b1 == std::string::npos || b2 == std::string::npos) { + return values; + } + std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); + std::string token; + while (std::getline(ss, token, ',')) { + try { + values.push_back(std::stoll(token)); + } catch (...) {} + } + return values; + }; + const auto segment_global_offsets = extract_int64_array("segment_global_offsets"); + const auto segment_local_offsets = extract_int64_array("segment_local_offsets"); + const auto segment_lengths = extract_int64_array("segment_lengths"); + CHECK_EQ(segment_global_offsets.size(), segment_local_offsets.size()); + CHECK_EQ(segment_global_offsets.size(), segment_lengths.size()); + for (size_t i = 0; i < segment_lengths.size(); ++i) { + entry.segments.push_back({.global_offset = segment_global_offsets[i], + .local_offset = segment_local_offsets[i], + .length = segment_lengths[i]}); + } + + auto ranks_pos = obj.find("\"stored_on_ranks\""); + if (ranks_pos != std::string::npos) { + auto b1 = obj.find('[', ranks_pos); + auto b2 = obj.find(']', b1); + std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); + std::string tok; + while (std::getline(ss, tok, ',')) { + try { + entry.stored_on_ranks.push_back(std::stoi(tok)); + } catch (...) {} + } + } + + meta.tensors.push_back(std::move(entry)); + obj_pos = obj_end + 1; + } + + LOG(INFO) << "[CKPT] Loaded metadata.json: " << meta.tensors.size() << " tensors, iteration=" << meta.iteration; + return meta; +} + +Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::path &checkpoint_dir) { + if (std::filesystem::exists(checkpoint_dir / "metadata.json")) { + return LoadSingleMetadata(checkpoint_dir); + } + + CheckpointMetadata merged; + for (const auto &entry : std::filesystem::directory_iterator(checkpoint_dir)) { + if (!entry.is_directory() || !entry.path().filename().string().starts_with("rank_") + || !std::filesystem::exists(entry.path() / "metadata.json")) { + continue; + } + auto rank_metadata = LoadSingleMetadata(entry.path()); + if (!rank_metadata.has_metadata) { + continue; + } + if (!merged.has_metadata) { + merged = rank_metadata; + merged.tensors.clear(); + } + for (auto &tensor : rank_metadata.tensors) { + tensor.file = (entry.path().filename() / tensor.file).generic_string(); + merged.tensors.push_back(std::move(tensor)); + } + } + LOG(INFO) << "[CKPT] Aggregated " << merged.tensors.size() << " tensor shards from rank manifests"; + return merged; +} + +void Checkpoint::SaveMetadataFile(const std::filesystem::path &path, const CheckpointMetadata &metadata) { + std::ofstream ofs(path); + CHECK(ofs.is_open()) << "Failed to write checkpoint metadata: " << path; + ofs << "{\n"; + ofs << " \"version\": 3,\n"; + ofs << " \"format\": \"infinitrain_sharded\",\n"; + ofs << " \"iteration\": " << metadata.iteration << ",\n"; + ofs << " \"parallel_config\": {\n"; + ofs << " \"tp_size\": " << metadata.parallel_config.tp_size << ",\n"; + ofs << " \"pp_size\": " << metadata.parallel_config.pp_size << ",\n"; + ofs << " \"dp_size\": " << metadata.parallel_config.dp_size << ",\n"; + ofs << " \"sp_size\": " << metadata.parallel_config.sp_size << "\n"; + ofs << " },\n"; + ofs << " \"tensors\": [\n"; + for (size_t i = 0; i < metadata.tensors.size(); ++i) { + const auto &tensor = metadata.tensors[i]; + ofs << " {\n"; + ofs << " \"key\": \"" << tensor.key << "\",\n"; + ofs << " \"dtype\": \"" << tensor.dtype_str << "\",\n"; + auto write_shape = [&](const char *name, const std::vector &shape) { + ofs << " \"" << name << "\": ["; + for (size_t d = 0; d < shape.size(); ++d) { ofs << shape[d] << (d + 1 < shape.size() ? ", " : ""); } + ofs << "],\n"; + }; + write_shape("global_shape", tensor.global_shape); + write_shape("local_shape", tensor.local_shape); + write_shape("global_offset", tensor.global_offset); + ofs << " \"axis_fragmentations\": ["; + for (size_t d = 0; d < tensor.axis_fragmentations.size(); ++d) { + ofs << tensor.axis_fragmentations[d] << (d + 1 < tensor.axis_fragmentations.size() ? ", " : ""); + } + ofs << "],\n"; + auto write_segments = [&](const char *name, auto member) { + ofs << " \"" << name << "\": ["; + for (size_t d = 0; d < tensor.segments.size(); ++d) { + ofs << tensor.segments[d].*member << (d + 1 < tensor.segments.size() ? ", " : ""); + } + ofs << "],\n"; + }; + write_segments("segment_global_offsets", &checkpoint::ShardSegment::global_offset); + write_segments("segment_local_offsets", &checkpoint::ShardSegment::local_offset); + write_segments("segment_lengths", &checkpoint::ShardSegment::length); + ofs << " \"file\": \"" << tensor.file << "\",\n"; + ofs << " \"offset\": " << tensor.offset << ",\n"; + ofs << " \"byte_size\": " << tensor.byte_size << ",\n"; + ofs << " \"pp_rank\": " << tensor.pp_rank << ",\n"; + ofs << " \"stored_on_ranks\": ["; + for (size_t r = 0; r < tensor.stored_on_ranks.size(); ++r) { + ofs << tensor.stored_on_ranks[r] << (r + 1 < tensor.stored_on_ranks.size() ? ", " : ""); + } + ofs << "]\n"; + ofs << " }" << (i + 1 < metadata.tensors.size() ? "," : "") << "\n"; + } + ofs << " ]\n"; + ofs << "}\n"; + CHECK(ofs.good()) << "Failed while writing checkpoint metadata: " << path; +} } // namespace infini_train diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index c6e31cddb..6bc771473 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -1,123 +1,230 @@ #include "infini_train/include/checkpoint/checkpoint_manager.h" -#include -#include +#include +#include #include #include +#include #include -#include -#include +#include #include #include "glog/logging.h" +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/reshard.h" +#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/lr_scheduler.h" +#include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/parallel_functional.h" +#include "infini_train/include/nn/parallel/work.h" #include "infini_train/include/tensor.h" using namespace infini_train; namespace nn = infini_train::nn; -// TODO(jym): ckpt is a new checkpoint format; bin is the legacy format. Keeping both as an interim solution; plan to -// consolidate into one later. -ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { - ResumeFromCheckpointResult result; - if (args.resume_root.empty()) { - LOG(INFO) << "No checkpoint specified for resume. Starting training from scratch."; - return result; +namespace { + +std::filesystem::path ResolveCheckpointDirectory(const std::filesystem::path &root) { + const auto latest_path = root / "latest_checkpointed_iteration.txt"; + if (!std::filesystem::exists(latest_path)) { + return root; } + std::ifstream latest(latest_path); + int64_t iteration = 0; + latest >> iteration; + const auto directory = root / std::format("iter_{:07d}", iteration); + CHECK(std::filesystem::exists(directory)) << "Latest checkpoint directory does not exist: " << directory; + return directory; +} - int ddp_world_size = nn::parallel::global::GetDataParallelSize(); - int tp_world_size = nn::parallel::global::GetTensorParallelSize(); - int sp_world_size = nn::parallel::global::GetSequenceParallelEnabled() ? tp_world_size : 1; - int pp_world_size = nn::parallel::global::GetPipelineParallelSize(); +void SynchronizeCheckpointRanks(const nn::Module &model) { + const auto parameters = model.Parameters(); + CHECK(!parameters.empty()) << "Cannot synchronize checkpoint save for a model without parameters"; + auto token = std::make_shared(std::vector{1}, DataType::kFLOAT32, parameters.front()->GetDevice()); + token->Fill(1.0f); + nn::parallel::function::AllReduce(token, nn::parallel::function::ReduceOpType::kSum, nullptr, true)->Synchronize(); +} - std::filesystem::path resume_dir = args.resume_root; - if (args.rank.IsParallel()) { - const auto rank_dir = resume_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); - if (std::filesystem::exists(rank_dir)) { - resume_dir = rank_dir; +void WaitForWriterManifests(const std::filesystem::path &staging_root, int tp_size, int pp_size, + int64_t expected_iteration) { + const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(10); + for (;;) { + bool ready = true; + for (int pp = 0; pp < pp_size && ready; ++pp) { + for (int tp = 0; tp < tp_size; ++tp) { + const int rank = nn::parallel::global::GetRankOf(0, tp, pp); + const auto manifest = staging_root / std::format("rank_{:06d}", rank) / "metadata.json"; + if (!std::filesystem::exists(manifest)) { + ready = false; + break; + } + const auto rank_metadata = Checkpoint::LoadMetadata(manifest.parent_path()); + if (!rank_metadata.has_metadata || rank_metadata.iteration != expected_iteration) { + ready = false; + break; + } + } + } + if (ready) { + return; } + CHECK(std::chrono::steady_clock::now() < deadline) + << "Timed out waiting for checkpoint manifests in " << staging_root; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); } +} - Checkpoint::Load(resume_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get()); +void WaitForGlobalMetadata(const std::filesystem::path &metadata_path) { + const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(10); + while (!std::filesystem::exists(metadata_path)) { + CHECK(std::chrono::steady_clock::now() < deadline) + << "Timed out waiting for global checkpoint metadata: " << metadata_path; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } +} - result.global_step = static_cast(args.state.global_step); +} // namespace - CHECK_EQ(args.state.n_layer, args.model_config.n_layer) - << "n_layer mismatch: ckpt=" << args.state.n_layer << ", config=" << args.model_config.n_layer; - CHECK_EQ(args.state.n_head, args.model_config.n_head) - << "n_head mismatch: ckpt=" << args.state.n_head << ", config=" << args.model_config.n_head; - CHECK_EQ(args.state.n_kv_head, args.model_config.n_kv_head) - << "n_kv_head mismatch: ckpt=" << args.state.n_kv_head << ", config=" << args.model_config.n_kv_head; - CHECK_EQ(args.state.n_embd, args.model_config.n_embd) - << "n_embd mismatch: ckpt=" << args.state.n_embd << ", config=" << args.model_config.n_embd; - CHECK_EQ(args.state.vocab_size, args.model_config.vocab_size) - << "vocab_size mismatch: ckpt=" << args.state.vocab_size << ", config=" << args.model_config.vocab_size; - - CHECK_EQ(args.state.ddp_size, ddp_world_size) << "DDP size mismatch: checkpoint has DDP=" << args.state.ddp_size - << ", but current run has DDP=" << ddp_world_size; - CHECK_EQ(args.state.tp_size, tp_world_size) - << "TP size mismatch: checkpoint has TP=" << args.state.tp_size << ", but current run has TP=" << tp_world_size; - CHECK_EQ(args.state.sp_size, sp_world_size) - << "SP size mismatch: checkpoint has SP=" << args.state.sp_size << ", but current run has SP=" << sp_world_size; - 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; +ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { + ResumeFromCheckpointResult result; + if (args.resume_root.empty()) { + LOG(INFO) << "No checkpoint specified for resume. Starting training from scratch."; + return result; + } + auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); + CHECK(std::filesystem::exists(checkpoint_dir / "metadata.json")) + << "Checkpoint metadata.json not found: " << checkpoint_dir; + auto metadata = Checkpoint::LoadMetadata(checkpoint_dir); + + CHECK(metadata.has_metadata); + CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + checkpoint::LoadDistributedCheckpoint(checkpoint_dir, *args.model, args.optimizer.get(), args.state, + args.lr_scheduler.get(), metadata); + + CHECK_EQ(args.state.n_layer, args.model_config.n_layer); + CHECK_EQ(args.state.n_head, args.model_config.n_head); + CHECK_EQ(args.state.n_kv_head, args.model_config.n_kv_head); + CHECK_EQ(args.state.n_embd, args.model_config.n_embd); + CHECK_GE(args.state.vocab_size, args.model_config.original_vocab_size) + << "Checkpoint vocabulary cannot represent the configured logical vocabulary"; + result.global_step = static_cast(args.state.global_step); 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_train_samples {}", args.state.global_step, args.state.consumed_train_samples); } - return result; } void SaveCheckpoint(const SaveCheckpointArgs &args) { - const auto ckpt_start = std::chrono::high_resolution_clock::now(); - - TrainerState state; - state.global_step = args.global_step; - 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; - state.n_embd = args.n_embd; - state.vocab_size = args.vocab_size; - state.ddp_size = args.ddp_size; - state.tp_size = args.tp_size; - state.sp_size = args.sp_size; - state.pp_size = args.pp_size; - - Checkpoint::Save(args.save_dir, args.model, args.optimizer, state, args.lr_scheduler); - - const auto ckpt_end = std::chrono::high_resolution_clock::now(); - const double ckpt_ms = std::chrono::duration(ckpt_end - ckpt_start).count(); - - if (!args.rank.IsMainRank()) { + const auto checkpoint_start = std::chrono::high_resolution_clock::now(); + TrainerState state{.global_step = args.global_step, + .consumed_train_samples = static_cast(args.consumed_train_samples), + .n_layer = args.n_layer, + .n_head = args.n_head, + .n_kv_head = args.n_kv_head, + .n_embd = args.n_embd, + .vocab_size = args.vocab_size, + .ddp_size = args.ddp_size, + .tp_size = args.tp_size, + .sp_size = args.sp_size, + .pp_size = args.pp_size}; + const auto iteration_dir = args.checkpoint_root_dir.empty() + ? args.save_dir + : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); + std::filesystem::create_directories(iteration_dir); + + const auto staging_root = iteration_dir / ".metadata_tmp"; + if (args.rank.IsMainRank()) { + std::filesystem::remove_all(staging_root); + } + SynchronizeCheckpointRanks(args.model); + + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(args.rank.GlobalRank(), dp_rank, tp_rank, pp_rank); + if (dp_rank != 0) { return; } - LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", args.save_dir.string(), ckpt_ms); + const auto rank_dir = iteration_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); + std::filesystem::create_directories(rank_dir); + auto sharded_state = args.model.ShardedStateDict(); + std::unordered_map> optimizer_state; + if (args.optimizer != nullptr) { + optimizer_state = args.optimizer->StateDict(); + auto optimizer_sharded_state = checkpoint::BuildOptimizerShardedStateDict(sharded_state, optimizer_state); + sharded_state.Merge(std::move(optimizer_sharded_state)); + } + auto write_items = checkpoint::SavePlanner::Plan(sharded_state, args.rank.GlobalRank()); + Checkpoint::SaveSharded(rank_dir, sharded_state, write_items, args.model.StateDict(), optimizer_state, state, + args.rank.GlobalRank()); + + const auto staging_rank_dir = staging_root / std::format("rank_{:06d}", args.rank.GlobalRank()); + std::filesystem::create_directories(staging_rank_dir); + const auto local_manifest = staging_rank_dir / "metadata.json"; + if (std::filesystem::exists(local_manifest)) { + std::filesystem::remove(local_manifest); + } + std::filesystem::rename(rank_dir / "metadata.json", local_manifest); - // FIXME(jym): Pruning currently relies on lexicographic sorting of directory names. - // This only works when step directories use zero-padded names (e.g. checkpoint_step_000042). - // If a future change introduces unpadded names, the prune order will be incorrect. - // Consider extracting the step number from the directory name and sorting numerically - // instead, once the checkpoint naming convention is finalized. - if (args.max_checkpoint_keep > 0 && std::filesystem::exists(args.checkpoint_root_dir)) { - std::vector ckpts; + if (args.rank.IsMainRank()) { + Checkpoint::SaveTrainerStateFile(iteration_dir / "trainer_state.json", state); + if (args.lr_scheduler != nullptr) { + Checkpoint::SaveLRSchedulerStateFile(iteration_dir / "lr_scheduler.ckpt", args.lr_scheduler->StateDict()); + } + WaitForWriterManifests(staging_root, args.tp_size, args.pp_size, args.global_step); + auto global_metadata = Checkpoint::LoadMetadata(staging_root); + CHECK(global_metadata.has_metadata); + const auto temporary_metadata = iteration_dir / "metadata.json.tmp"; + const auto final_metadata = iteration_dir / "metadata.json"; + if (std::filesystem::exists(temporary_metadata)) { + std::filesystem::remove(temporary_metadata); + } + Checkpoint::SaveMetadataFile(temporary_metadata, global_metadata); + if (std::filesystem::exists(final_metadata)) { + std::filesystem::remove(final_metadata); + } + std::filesystem::rename(temporary_metadata, final_metadata); + std::filesystem::remove_all(staging_root); + } else { + WaitForGlobalMetadata(iteration_dir / "metadata.json"); + } + + if (args.rank.IsMainRank() && !args.checkpoint_root_dir.empty()) { + const auto latest = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt"; + const auto temporary_latest = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt.tmp"; + { + std::ofstream output(temporary_latest); + CHECK(output.is_open()); + output << args.global_step; + } + if (std::filesystem::exists(latest)) { + std::filesystem::remove(latest); + } + std::filesystem::rename(temporary_latest, latest); + } + + if (args.rank.IsMainRank() && args.max_checkpoint_keep > 0 && std::filesystem::exists(args.checkpoint_root_dir)) { + std::vector checkpoints; for (const auto &entry : std::filesystem::directory_iterator(args.checkpoint_root_dir)) { - if (entry.is_directory() && entry.path().filename().string().starts_with("checkpoint_step_")) { - ckpts.push_back(entry.path()); + if (entry.is_directory() && entry.path().filename().string().starts_with("iter_")) { + checkpoints.push_back(entry.path()); } } - std::sort(ckpts.begin(), ckpts.end()); - while (ckpts.size() > args.max_checkpoint_keep) { - std::filesystem::remove_all(ckpts.front()); - ckpts.erase(ckpts.begin()); + std::sort(checkpoints.begin(), checkpoints.end()); + while (checkpoints.size() > args.max_checkpoint_keep) { + std::filesystem::remove_all(checkpoints.front()); + checkpoints.erase(checkpoints.begin()); } } + + const auto checkpoint_end = std::chrono::high_resolution_clock::now(); + const double elapsed_ms = std::chrono::duration(checkpoint_end - checkpoint_start).count(); + LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", iteration_dir.string(), elapsed_ms); } 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/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc new file mode 100644 index 000000000..871070d19 --- /dev/null +++ b/infini_train/src/checkpoint/load_planner.cc @@ -0,0 +1,244 @@ +#include "infini_train/include/checkpoint/load_planner.h" + +#include +#include +#include +#include + +#include "glog/logging.h" + +namespace infini_train::checkpoint { +namespace { + +DataType StringToDataType(const std::string &value) { + static const std::unordered_map legacy_names = { + {"bfloat16", DataType::kBFLOAT16}, + {"float16", DataType::kFLOAT16}, + {"float32", DataType::kFLOAT32}, + {"float64", DataType::kFLOAT64}, + }; + if (const auto it = legacy_names.find(value); it != legacy_names.end()) { + return it->second; + } + for (const auto &[dtype, description] : kDataTypeToDesc) { + if (description == value) { + return dtype; + } + } + LOG(FATAL) << "Unsupported checkpoint tensor dtype: " << value; + return DataType::kFLOAT32; +} + +int FragmentedAxis(const std::vector &axis_fragmentations) { + int fragmented_axis = -1; + for (size_t dim = 0; dim < axis_fragmentations.size(); ++dim) { + if (axis_fragmentations[dim] <= 1) { + continue; + } + CHECK_EQ(fragmented_axis, -1) << "Multi-axis checkpoint sharding is not supported yet"; + fragmented_axis = static_cast(dim); + } + return fragmented_axis; +} + +bool IsVocabularyTensor(const std::string &key) { + std::string parameter_key = key; + if (parameter_key.starts_with("adam.m.")) { + parameter_key = parameter_key.substr(7); + } else if (parameter_key.starts_with("adam.v.")) { + parameter_key = parameter_key.substr(7); + } + return parameter_key == "transformer.wte.weight" || parameter_key == "lm_head.weight"; +} + +bool IsPaddingCompatible(const std::string &key, const std::vector &source, + const std::vector &target) { + if (!IsVocabularyTensor(key) || source.size() != target.size() || source.empty()) { + return false; + } + for (size_t dim = 1; dim < source.size(); ++dim) { + if (source[dim] != target[dim]) { + return false; + } + } + return true; +} + +void ValidateCoordinates(const std::string &key, const std::vector &global_shape, + const std::vector &local_shape, const std::vector &global_offset, + const std::vector &axis_fragmentations) { + CHECK_EQ(global_shape.size(), local_shape.size()) << "Invalid local rank for tensor " << key; + CHECK_EQ(global_shape.size(), global_offset.size()) << "Invalid offset rank for tensor " << key; + CHECK_EQ(global_shape.size(), axis_fragmentations.size()) << "Invalid fragmentation rank for tensor " << key; + for (size_t dim = 0; dim < global_shape.size(); ++dim) { + CHECK_GT(global_shape[dim], 0) << "Invalid global shape for tensor " << key; + CHECK_GT(local_shape[dim], 0) << "Invalid local shape for tensor " << key; + CHECK_GE(global_offset[dim], 0) << "Invalid global offset for tensor " << key; + CHECK_LE(global_offset[dim] + local_shape[dim], global_shape[dim]) + << "Shard exceeds global shape for tensor " << key; + CHECK_GE(axis_fragmentations[dim], 1) << "Invalid axis fragmentation for tensor " << key; + } +} + +} // namespace + +LoadPlan LoadPlanner::PlanReshard(const Checkpoint::CheckpointMetadata &metadata, + const ShardedStateDict &target_state_dict) { + LoadPlan plan; + for (const auto &[key, target] : target_state_dict.tensors) { + ValidateCoordinates(key, target.global_shape, target.local_shape, target.global_offset, + target.axis_fragmentations); + TargetTensorPlan tensor_plan{.key = key, + .dtype = target.dtype, + .global_shape = target.global_shape, + .target_shape = target.local_shape, + .shard_dim = FragmentedAxis(target.axis_fragmentations)}; + + std::vector candidates; + for (const auto &entry : metadata.tensors) { + if (entry.key == key) { + candidates.push_back(&entry); + } + } + CHECK(!candidates.empty()) << "No saved shard found for target tensor: " << key; + const int saved_axis = FragmentedAxis(candidates.front()->axis_fragmentations); + for (const auto *source : candidates) { + ValidateCoordinates(key, source->global_shape, source->local_shape, source->global_offset, + source->axis_fragmentations); + CHECK(source->global_shape == target.global_shape + || IsPaddingCompatible(key, source->global_shape, target.global_shape)) + << "Global shape changed for tensor " << key; + CHECK_EQ(FragmentedAxis(source->axis_fragmentations), saved_axis) + << "Inconsistent saved shard dimensions for tensor " << key; + } + + if (!target.segments.empty()) { + if (tensor_plan.shard_dim < 0) { + tensor_plan.shard_dim = 0; + } + int64_t target_covered = 0; + for (const auto &target_segment : target.segments) { + CHECK_EQ(target_segment.local_offset, target_covered) + << "Gap or overlap in target segments for " << key; + CHECK_GT(target_segment.length, 0); + CHECK_LE(target_segment.global_offset + target_segment.length, + target.global_shape[tensor_plan.shard_dim]); + target_covered += target_segment.length; + for (const auto *source : candidates) { + CHECK(!source->segments.empty()) << "Saved checkpoint lacks segmented layout metadata for " << key; + for (const auto &source_segment : source->segments) { + const int64_t overlap_start + = std::max(source_segment.global_offset, target_segment.global_offset); + const int64_t overlap_end = std::min(source_segment.global_offset + source_segment.length, + target_segment.global_offset + target_segment.length); + if (overlap_start >= overlap_end) { + continue; + } + tensor_plan.reads.push_back({ + .key = key, + .filename = source->file, + .dtype = StringToDataType(source->dtype_str), + .global_shape = source->global_shape, + .byte_size = source->byte_size, + .data_offset = source->offset, + .shard_dim = tensor_plan.shard_dim, + .source_offset = source_segment.local_offset + overlap_start - source_segment.global_offset, + .target_offset = target_segment.local_offset + overlap_start - target_segment.global_offset, + .length = overlap_end - overlap_start, + .source_shape = source->local_shape, + }); + } + } + } + CHECK_EQ(target_covered, target.local_shape[tensor_plan.shard_dim]) + << "Segmented layout does not cover target local tensor " << key; + std::sort( + tensor_plan.reads.begin(), tensor_plan.reads.end(), + [](const ReadItem &left, const ReadItem &right) { return left.target_offset < right.target_offset; }); + int64_t covered = 0; + for (const auto &read : tensor_plan.reads) { + CHECK_EQ(read.target_offset, covered) << "Gap or overlap in segmented target plan for " << key; + covered += read.length; + } + CHECK_EQ(covered, target_covered) << "Incomplete segmented target plan for " << key; + plan.tensors.emplace(key, std::move(tensor_plan)); + continue; + } + + if (tensor_plan.shard_dim < 0) { + tensor_plan.shard_dim = saved_axis; + } + if (tensor_plan.shard_dim < 0 && candidates.front()->global_shape != target.global_shape) { + tensor_plan.shard_dim = 0; + } + if (saved_axis >= 0 && FragmentedAxis(target.axis_fragmentations) >= 0) { + CHECK_EQ(saved_axis, tensor_plan.shard_dim) << "Shard dimension changed for tensor " << key; + } + + if (tensor_plan.shard_dim < 0) { + const auto *source = candidates.front(); + tensor_plan.reads.push_back({.key = key, + .filename = source->file, + .dtype = StringToDataType(source->dtype_str), + .global_shape = source->global_shape, + .byte_size = source->byte_size, + .data_offset = source->offset, + .shard_dim = -1, + .source_shape = source->local_shape}); + plan.tensors.emplace(key, std::move(tensor_plan)); + continue; + } + + const int dim = tensor_plan.shard_dim; + const int64_t target_start = target.global_offset[dim]; + const int64_t target_length = target.local_shape[dim]; + const int64_t target_end = target_start + target_length; + std::set> seen_source_ranges; + + for (const auto *source : candidates) { + const int64_t saved_start = source->global_offset[dim]; + const int64_t saved_length = source->local_shape[dim]; + if (!seen_source_ranges.emplace(saved_start, saved_length).second) { + continue; + } + const int64_t overlap_start = std::max(saved_start, target_start); + const int64_t overlap_end = std::min(saved_start + saved_length, target_end); + if (overlap_start >= overlap_end) { + continue; + } + + tensor_plan.reads.push_back({.key = key, + .filename = source->file, + .dtype = StringToDataType(source->dtype_str), + .global_shape = source->global_shape, + .byte_size = source->byte_size, + .data_offset = source->offset, + .shard_dim = dim, + .source_offset = overlap_start - saved_start, + .target_offset = overlap_start - target_start, + .length = overlap_end - overlap_start, + .source_shape = source->local_shape}); + } + + std::sort(tensor_plan.reads.begin(), tensor_plan.reads.end(), + [](const ReadItem &left, const ReadItem &right) { return left.target_offset < right.target_offset; }); + int64_t covered = 0; + for (const auto &read : tensor_plan.reads) { + CHECK_EQ(read.target_offset, covered) << "Gap or overlap in target shard plan for " << key; + covered += read.length; + } + if (covered < target_length) { + CHECK(IsVocabularyTensor(key)) << "Incomplete target shard plan for " << key; + CHECK_EQ(dim, 0) << "Vocabulary padding is only supported along dim 0"; + CHECK_EQ(target_start + covered, candidates.front()->global_shape[0]) + << "Only trailing vocabulary padding is supported for " << key; + tensor_plan.trailing_zero_fill = target_length - covered; + covered = target_length; + } + CHECK_EQ(covered, target_length) << "Incomplete target shard plan for " << key; + plan.tensors.emplace(key, std::move(tensor_plan)); + } + return plan; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/checkpoint/load_strategy.cc b/infini_train/src/checkpoint/load_strategy.cc new file mode 100644 index 000000000..476637881 --- /dev/null +++ b/infini_train/src/checkpoint/load_strategy.cc @@ -0,0 +1,129 @@ +#include "infini_train/include/checkpoint/load_strategy.h" + +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/nn/functional.h" +#include "infini_train/include/tensor.h" +#include "infini_train/include/utils/string_utils.h" + +namespace infini_train::checkpoint { +namespace { + +using FileCache = std::unordered_map>; + +std::ifstream &GetFile(FileCache &cache, const std::filesystem::path &checkpoint_dir, const std::string &filename) { + auto it = cache.find(filename); + if (it == cache.end()) { + auto stream = std::make_unique(checkpoint_dir / filename, std::ios::binary); + CHECK(stream->is_open()) << "Failed to open checkpoint file: " << checkpoint_dir / filename; + it = cache.emplace(filename, std::move(stream)).first; + } + return *it->second; +} + +void ReadAt(std::ifstream &stream, uint64_t offset, void *destination, uint64_t byte_size, const std::string &key, + const std::string &filename) { + stream.clear(); + stream.seekg(static_cast(offset), std::ios::beg); + CHECK(stream.good()) << "Failed to seek tensor " << key << " in " << filename; + stream.read(static_cast(destination), static_cast(byte_size)); + CHECK_EQ(static_cast(stream.gcount()), byte_size) << "Truncated tensor " << key << " in " << filename; +} + +std::shared_ptr ReadTensor(std::ifstream &stream, const ReadItem &read) { + auto tensor = std::make_shared(read.source_shape, read.dtype, Device()); + CHECK_EQ(read.byte_size, tensor->SizeInBytes()) + << "Tensor byte size mismatch for " << read.key << " in " << read.filename; + ReadAt(stream, read.data_offset, tensor->DataPtr(), read.byte_size, read.key, read.filename); + return tensor; +} + +std::shared_ptr ReadTensorRegion(std::ifstream &stream, const ReadItem &read) { + CHECK_GE(read.shard_dim, 0); + CHECK_LT(read.shard_dim, static_cast(read.source_shape.size())); + CHECK_GE(read.source_offset, 0); + CHECK_GT(read.length, 0); + CHECK_LE(read.source_offset + read.length, read.source_shape[read.shard_dim]); + + uint64_t source_numel = 1; + for (const auto size : read.source_shape) { + CHECK_GT(size, 0); + source_numel *= static_cast(size); + } + CHECK_EQ(read.byte_size % source_numel, 0u) + << "Invalid tensor byte size for " << read.key << " in " << read.filename; + const uint64_t element_size = read.byte_size / source_numel; + + uint64_t outer = 1; + for (int dim = 0; dim < read.shard_dim; ++dim) { outer *= static_cast(read.source_shape[dim]); } + uint64_t inner = 1; + for (size_t dim = static_cast(read.shard_dim + 1); dim < read.source_shape.size(); ++dim) { + inner *= static_cast(read.source_shape[dim]); + } + + auto region_shape = read.source_shape; + region_shape[read.shard_dim] = read.length; + auto tensor = std::make_shared(region_shape, read.dtype, Device()); + const uint64_t block_bytes = static_cast(read.length) * inner * element_size; + const uint64_t source_stride_bytes + = static_cast(read.source_shape[read.shard_dim]) * inner * element_size; + const uint64_t first_block_offset + = read.data_offset + static_cast(read.source_offset) * inner * element_size; + CHECK_EQ(outer * block_bytes, tensor->SizeInBytes()); + + auto *destination = static_cast(tensor->DataPtr()); + for (uint64_t block = 0; block < outer; ++block) { + ReadAt(stream, first_block_offset + block * source_stride_bytes, destination + block * block_bytes, block_bytes, + read.key, read.filename); + } + return tensor; +} + +} // namespace + +LoadedStateDict IndexedRegionLoadStrategy::Execute(const std::filesystem::path &checkpoint_dir, const LoadPlan &plan) { + FileCache file_cache; + LoadedStateDict result; + + for (const auto &[key, tensor_plan] : plan.tensors) { + CHECK(!tensor_plan.reads.empty() + || (tensor_plan.shard_dim >= 0 + && tensor_plan.trailing_zero_fill == tensor_plan.target_shape[tensor_plan.shard_dim])) + << "No reads or padding planned for target tensor: " << key; + std::vector> pieces; + pieces.reserve(tensor_plan.reads.size() + (tensor_plan.trailing_zero_fill > 0 ? 1 : 0)); + + for (const auto &read : tensor_plan.reads) { + CHECK_GT(read.data_offset, 0) << "Checkpoint metadata lacks a valid tensor data offset for " << key + << "; regenerate the checkpoint with the current format"; + auto &stream = GetFile(file_cache, checkpoint_dir, read.filename); + pieces.push_back(read.shard_dim < 0 ? ReadTensor(stream, read) : ReadTensorRegion(stream, read)); + } + + if (tensor_plan.trailing_zero_fill > 0) { + CHECK_GE(tensor_plan.shard_dim, 0); + auto padding_shape = tensor_plan.target_shape; + padding_shape[tensor_plan.shard_dim] = tensor_plan.trailing_zero_fill; + auto padding = std::make_shared(padding_shape, tensor_plan.dtype, Device()); + padding->Fill(0.0f); + pieces.push_back(std::move(padding)); + } + + auto target = pieces.front(); + if (pieces.size() > 1) { + CHECK_GE(tensor_plan.shard_dim, 0); + target = nn::function::Concat(pieces, tensor_plan.shard_dim)->Contiguous(); + } + CHECK(target->Dims() == tensor_plan.target_shape) + << "Target shard shape mismatch for " << key + << ": expected=" << utils::DimsToString(tensor_plan.target_shape) + << ", got=" << utils::DimsToString(target->Dims()); + result.emplace(key, std::move(target)); + } + return result; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/checkpoint/reshard.cc b/infini_train/src/checkpoint/reshard.cc new file mode 100644 index 000000000..fb8842f8a --- /dev/null +++ b/infini_train/src/checkpoint/reshard.cc @@ -0,0 +1,71 @@ +#include "infini_train/include/checkpoint/reshard.h" + +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/load_strategy.h" +#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/lr_scheduler.h" +#include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/optimizer.h" + +namespace infini_train::checkpoint { + +void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, LRScheduler *lr_scheduler, + const Checkpoint::CheckpointMetadata &metadata) { + CHECK(metadata.has_metadata); + CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + auto model_sharded_state = model.ShardedStateDict(); + auto plan = LoadPlanner::PlanReshard(metadata, model_sharded_state); + IndexedRegionLoadStrategy strategy; + auto result = strategy.Execute(checkpoint_dir, plan); + model.LoadStateDict(result); + + state = Checkpoint::LoadTrainerStateFile(checkpoint_dir / "trainer_state.json"); + const int current_tp = nn::parallel::global::GetTensorParallelSize(); + const int current_pp = nn::parallel::global::GetPipelineParallelSize(); + const bool topology_changed + = current_tp != metadata.parallel_config.tp_size || current_pp != metadata.parallel_config.pp_size; + state.tp_size = current_tp; + state.pp_size = current_pp; + state.ddp_size = nn::parallel::global::GetDataParallelSize(); + state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; + + if (optimizer != nullptr) { + if (topology_changed) { + const auto initialized_optimizer_state = optimizer->StateDict(); + auto optimizer_sharded_state + = BuildOptimizerShardedStateDict(model_sharded_state, initialized_optimizer_state); + auto optimizer_plan = LoadPlanner::PlanReshard(metadata, optimizer_sharded_state); + auto loaded_optimizer_state = strategy.Execute(checkpoint_dir, optimizer_plan); + optimizer->LoadStateDict(loaded_optimizer_state); + LOG(INFO) << "[CKPT] Resharded " << loaded_optimizer_state.size() + << " optimizer tensors across TP/PP topology change"; + } else { + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(nn::parallel::global::thread_global_rank, dp_rank, tp_rank, pp_rank); + const int writer_rank = nn::parallel::global::GetRankOf(0, tp_rank, pp_rank); + const auto optimizer_path = checkpoint_dir / std::format("rank_{:06d}/optimizer.ckpt", writer_rank); + CHECK(std::filesystem::exists(optimizer_path)) + << "Optimizer checkpoint not found for current_rank=" << nn::parallel::global::thread_global_rank + << ", coords=(dp=" << dp_rank << ", tp=" << tp_rank << ", pp=" << pp_rank + << "), writer_rank=" << writer_rank << ": " << optimizer_path; + LOG(INFO) << "[CKPT] Loading optimizer for current_rank=" << nn::parallel::global::thread_global_rank + << " from writer_rank=" << writer_rank << ": " << optimizer_path; + optimizer->LoadStateDict(Checkpoint::LoadStateDictFile(optimizer_path)); + } + } + if (lr_scheduler != nullptr && std::filesystem::exists(checkpoint_dir / "lr_scheduler.ckpt")) { + lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(checkpoint_dir / "lr_scheduler.ckpt")); + } + LOG(INFO) << "[CKPT] Restored " << result.size() + << " tensors with overlap reads from TP=" << metadata.parallel_config.tp_size + << ", PP=" << metadata.parallel_config.pp_size << " to TP=" << current_tp << ", PP=" << current_pp; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/checkpoint/save_planner.cc b/infini_train/src/checkpoint/save_planner.cc new file mode 100644 index 000000000..17ae1a85b --- /dev/null +++ b/infini_train/src/checkpoint/save_planner.cc @@ -0,0 +1,75 @@ +#include "infini_train/include/checkpoint/save_planner.h" + +#include "glog/logging.h" + +#include "infini_train/include/tensor.h" + +namespace infini_train::checkpoint { + +ShardedStateDict +BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, + const std::unordered_map> &optimizer_state) { + ShardedStateDict result; + for (const auto &[key, tensor] : optimizer_state) { + if (key == "adam.t") { + ShardedTensor info; + info.key = key; + info.local_key = key; + info.dtype = tensor->Dtype(); + info.global_shape = tensor->Dims(); + info.local_shape = tensor->Dims(); + info.global_offset.assign(tensor->Dims().size(), 0); + info.axis_fragmentations.assign(tensor->Dims().size(), 1); + result.tensors.emplace(key, std::move(info)); + continue; + } + + std::string parameter_key; + if (key.starts_with("adam.m.")) { + parameter_key = key.substr(std::string("adam.m.").size()); + } else if (key.starts_with("adam.v.")) { + parameter_key = key.substr(std::string("adam.v.").size()); + } else { + CHECK(false) << "Unsupported optimizer state key: " << key; + } + + auto model_it = model_state.tensors.find(parameter_key); + CHECK(model_it != model_state.tensors.end()) + << "Optimizer state " << key << " has no matching named model parameter. " + << "Optimizer resharding requires set_parameter_names()."; + auto info = model_it->second; + info.key = key; + info.local_key = key; + result.tensors.emplace(key, std::move(info)); + } + return result; +} + +std::vector SavePlanner::Plan(const ShardedStateDict &sd, int rank) { + std::vector items; + uint64_t model_offset = 0; + uint64_t optim_offset = 0; + + for (auto &[key, info] : sd.tensors) { + bool is_optimizer = key.starts_with("adam."); + uint64_t &offset = is_optimizer ? optim_offset : model_offset; + + WriteItem item; + item.key = key; + item.filename = is_optimizer ? "optimizer.ckpt" : "model.ckpt"; + item.offset = offset; + item.byte_size = TensorByteSize(info.dtype, info.local_shape); + item.dtype = info.dtype; + item.local_shape = info.local_shape; + item.global_offset = info.global_offset; + item.axis_fragmentations = info.axis_fragmentations; + item.rank = rank; + + items.push_back(std::move(item)); + offset += items.back().byte_size; + } + + return items; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 498bcc075..ab8fd6ff8 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -188,6 +188,44 @@ std::unordered_map> Module::StateDict() con return state; } +checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix) const { + checkpoint::ShardedStateDict sd; + + for (auto &[name, param] : parameters_) { + checkpoint::ShardedTensor info; + info.key = prefix.empty() ? name : prefix + "." + name; + info.dtype = param->Dtype(); + info.global_shape = param->Dims(); + info.local_shape = param->Dims(); + info.global_offset.assign(param->Dims().size(), 0); + info.axis_fragmentations.assign(param->Dims().size(), 1); + sd.tensors[info.key] = std::move(info); + } + + for (auto &[name, buffer] : buffers_) { + checkpoint::ShardedTensor info; + info.key = prefix.empty() ? name : prefix + "." + name; + info.dtype = buffer->Dtype(); + info.global_shape = buffer->Dims(); + info.local_shape = buffer->Dims(); + info.global_offset.assign(buffer->Dims().size(), 0); + info.axis_fragmentations.assign(buffer->Dims().size(), 1); + sd.tensors[info.key] = std::move(info); + } + + for (auto &[name, module] : modules_) { + if (name.starts_with("__pp")) { + continue; + } + + auto child_prefix = prefix.empty() ? name : prefix + "." + name; + auto child_sd = module->ShardedStateDict(child_prefix); + sd.Merge(std::move(child_sd)); + } + + return sd; +} + void Module::LoadStateDict(const std::unordered_map> &state_dict) { // Stage 1: Validate all keys, shapes, and dtypes without copying std::vector error_msgs; diff --git a/infini_train/src/nn/modules/transformer/causal_self_attention.cc b/infini_train/src/nn/modules/transformer/causal_self_attention.cc index 9c32fa23d..5db2ed1d9 100644 --- a/infini_train/src/nn/modules/transformer/causal_self_attention.cc +++ b/infini_train/src/nn/modules/transformer/causal_self_attention.cc @@ -81,6 +81,40 @@ void CausalSelfAttention::SetupAttention(const TransformerConfig &config) { } } +checkpoint::ShardedStateDict CausalSelfAttention::ShardedStateDict(const std::string &prefix) const { + auto state = Module::ShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + const int rank = parallel::tp_rank; + const int64_t q_global = n_head_ * head_dim_; + const int64_t kv_global = n_kv_head_ * head_dim_; + const int64_t q_local = q_global / tp_size; + const int64_t kv_local = kv_global / tp_size; + + const auto c_attn_prefix = prefix.empty() ? kCAttnLayerName : prefix + "." + kCAttnLayerName; + auto set_qkv_segments = [&](const std::string ¶meter_name) { + const auto key = c_attn_prefix + "." + parameter_name; + auto it = state.tensors.find(key); + if (it == state.tensors.end()) { + return; + } + auto &tensor = it->second; + tensor.global_offset.assign(tensor.global_shape.size(), 0); + tensor.segments = { + {.global_offset = rank * q_local, .local_offset = 0, .length = q_local}, + {.global_offset = q_global + rank * kv_local, .local_offset = q_local, .length = kv_local}, + {.global_offset = q_global + kv_global + rank * kv_local, + .local_offset = q_local + kv_local, + .length = kv_local}, + }; + }; + + set_qkv_segments(parallel::ColumnParallelLinear::kParamWeightName); + if (config_.add_bias_linear) { + set_qkv_segments(parallel::ColumnParallelLinear::kParamBiasName); + } + return state; +} + std::shared_ptr CausalSelfAttention::RepeatKV(const std::shared_ptr &x, int64_t n_rep) { const auto &shape = x->Dims(); diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index 99a739d2d..d6e6e28a0 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -272,6 +272,83 @@ TransformerModel::TransformerModel(const TransformerConfig config) } } +namespace { + +std::vector GlobalLayerIndices(const parallel::StageInfo &stage_info) { + std::vector indices; + for (const auto &[start, end] : stage_info.layer_ranges_per_chunk) { + for (int layer = start; layer < end; ++layer) { indices.push_back(layer); } + } + std::sort(indices.begin(), indices.end()); + return indices; +} + +std::string RemapLayerKey(const std::string &key, const std::vector &from, const std::vector &to) { + const std::string marker + = std::string(TransformerModel::kTransformerModelName) + "." + TransformerChunk::kHLayerName + "."; + const auto marker_pos = key.find(marker); + if (marker_pos == std::string::npos) { + return key; + } + const auto index_start = marker_pos + marker.size(); + const auto index_end = key.find('.', index_start); + if (index_end == std::string::npos) { + return key; + } + int layer = -1; + try { + layer = std::stoi(key.substr(index_start, index_end - index_start)); + } catch (...) { return key; } + const auto it = std::find(from.begin(), from.end(), layer); + if (it == from.end()) { + return key; + } + const auto mapped = to[static_cast(std::distance(from.begin(), it))]; + return key.substr(0, index_start) + std::to_string(mapped) + key.substr(index_end); +} + +} // namespace + +checkpoint::ShardedStateDict TransformerModel::ShardedStateDict(const std::string &prefix) const { + auto local_state = Module::ShardedStateDict(prefix); + const auto global_layers = GlobalLayerIndices(stage_info_); + std::vector local_layers(global_layers.size()); + std::iota(local_layers.begin(), local_layers.end(), 0); + + checkpoint::ShardedStateDict global_state; + for (auto &[local_key, tensor] : local_state.tensors) { + const auto global_key = RemapLayerKey(local_key, local_layers, global_layers); + if (global_key != local_key) { + tensor.local_key = local_key; + tensor.key = global_key; + } + global_state.tensors.emplace(global_key, std::move(tensor)); + } + return global_state; +} + +std::vector>> +TransformerModel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + auto parameters = Module::NamedParameters(prefix, recurse, remove_duplicate); + const auto global_layers = GlobalLayerIndices(stage_info_); + std::vector local_layers(global_layers.size()); + std::iota(local_layers.begin(), local_layers.end(), 0); + for (auto &[name, parameter] : parameters) { name = RemapLayerKey(name, local_layers, global_layers); } + return parameters; +} + +void TransformerModel::LoadStateDict(const std::unordered_map> &state_dict) { + const auto global_layers = GlobalLayerIndices(stage_info_); + std::vector local_layers(global_layers.size()); + std::iota(local_layers.begin(), local_layers.end(), 0); + + std::unordered_map> local_state; + for (const auto &[global_key, tensor] : state_dict) { + local_state.emplace(RemapLayerKey(global_key, global_layers, local_layers), tensor); + } + Module::LoadStateDict(local_state); +} + std::vector> TransformerModel::Forward(const std::vector> &x) { auto x1 = (*modules_[kPPFirstStageName])(x); for (int chunk_idx = 0; chunk_idx < stage_info_.layer_ranges_per_chunk.size(); ++chunk_idx) { diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index dd05a8d71..100f5e5aa 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -211,6 +211,24 @@ void DistributedDataParallel::OnGradReady(const std::shared_ptr ¶m) } } +std::vector>> +DistributedDataParallel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + return modules_.at(kModuleName)->NamedParameters(prefix, recurse, remove_duplicate); +} + +std::unordered_map> DistributedDataParallel::StateDict() const { + return modules_.at(kModuleName)->StateDict(); +} + +checkpoint::ShardedStateDict DistributedDataParallel::ShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->ShardedStateDict(prefix); +} + +void DistributedDataParallel::LoadStateDict( + const std::unordered_map> &state_dict) { + modules_.at(kModuleName)->LoadStateDict(state_dict); +} + std::vector> DistributedDataParallel::Forward(const std::vector> &input_tensors) { auto outputs = (*modules_[kModuleName])(input_tensors); diff --git a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc index c0369cdeb..0dc34fdfb 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc @@ -104,4 +104,21 @@ PipelineParallel::PipelineParallel(const std::shared_ptr module, int num } std::vector> *PipelineParallel::mutable_chunks() { return pipeline_stage_->mutable_chunks(); } + +std::vector>> +PipelineParallel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + return modules_.at(kModuleName)->NamedParameters(prefix, recurse, remove_duplicate); +} + +std::unordered_map> PipelineParallel::StateDict() const { + return modules_.at(kModuleName)->StateDict(); +} + +checkpoint::ShardedStateDict PipelineParallel::ShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->ShardedStateDict(prefix); +} + +void PipelineParallel::LoadStateDict(const std::unordered_map> &state_dict) { + modules_.at(kModuleName)->LoadStateDict(state_dict); +} } // namespace infini_train::nn::parallel diff --git a/infini_train/src/nn/parallel/tensor_parallel.cc b/infini_train/src/nn/parallel/tensor_parallel.cc index b16c526e6..b5a54b8d9 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -284,6 +284,36 @@ bool ColumnParallelLinear::input_is_parallel() const { return input_is_parallel_ bool ColumnParallelLinear::skip_bias_add() const { return skip_bias_add_; } bool ColumnParallelLinear::sequence_parallel() const { return sequence_parallel_; } +checkpoint::ShardedStateDict ColumnParallelLinear::ShardedStateDict(const std::string &prefix) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {output_size_per_partition_ * tp_size, weight->Dims()[1]}; + w.local_shape = weight->Dims(); + w.global_offset = {output_size_per_partition_ * tp_rank, 0}; + w.axis_fragmentations = {tp_size, 1}; + sd.tensors[w.key] = std::move(w); + + // Bias is also split along dim=0 + if (bias_) { + auto &bias = parameter(kParamBiasName); + checkpoint::ShardedTensor b; + b.key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; + b.dtype = bias->Dtype(); + b.global_shape = {static_cast(output_size_per_partition_ * tp_size)}; + b.local_shape = bias->Dims(); + b.global_offset = {output_size_per_partition_ * tp_rank}; + b.axis_fragmentations = {tp_size}; + sd.tensors[b.key] = std::move(b); + } + + return sd; +} + RowParallelLinear::RowParallelLinear(int64_t in_features, int64_t out_features, bool bias, bool reduce_output, bool input_is_parallel, bool skip_bias_add, bool sequence_parallel) : CloneableModule(kType), bias_(bias), reduce_output_(reduce_output), input_is_parallel_(input_is_parallel), @@ -339,6 +369,36 @@ bool RowParallelLinear::input_is_parallel() const { return input_is_parallel_; } bool RowParallelLinear::skip_bias_add() const { return skip_bias_add_; } bool RowParallelLinear::sequence_parallel() const { return sequence_parallel_; } +checkpoint::ShardedStateDict RowParallelLinear::ShardedStateDict(const std::string &prefix) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {weight->Dims()[0], input_size_per_partition_ * tp_size}; + w.local_shape = weight->Dims(); + w.global_offset = {0, input_size_per_partition_ * tp_rank}; + w.axis_fragmentations = {1, tp_size}; + sd.tensors[w.key] = std::move(w); + + // Bias is NOT sharded in RowParallelLinear + if (bias_) { + auto &bias = parameter(kParamBiasName); + checkpoint::ShardedTensor b; + b.key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; + b.dtype = bias->Dtype(); + b.global_shape = bias->Dims(); + b.local_shape = bias->Dims(); + b.global_offset = {0}; + b.axis_fragmentations = {1}; + sd.tensors[b.key] = std::move(b); + } + + return sd; +} + VocabParallelEmbedding::VocabParallelEmbedding(int64_t num_embeddings, int64_t embedding_dim, bool reduce_scatter_embeddings) : CloneableModule(kType), embedding_dim_(embedding_dim), reduce_scatter_embeddings_(reduce_scatter_embeddings) { @@ -395,6 +455,23 @@ VocabParallelEmbedding::Forward(const std::vector> &inpu return {output}; } +checkpoint::ShardedStateDict VocabParallelEmbedding::ShardedStateDict(const std::string &prefix) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {vocab_size_global_, embedding_dim_}; + w.local_shape = weight->Dims(); + w.global_offset = {vocab_start_index_, 0}; + w.axis_fragmentations = {tp_size, 1}; + sd.tensors[w.key] = std::move(w); + + return sd; +} + std::vector> VocabParallelCrossEntropy::Forward(const std::vector> &input_tensors) { CHECK_EQ(input_tensors.size(), 2) << kType << " expects {logits, target}"; @@ -465,7 +542,7 @@ VocabParallelCrossEntropy::Forward(const std::vector> &i auto sum_exp_local = exp_local->Sum(-1); auto sum_exp = (tp_size > 1) ? ReduceFromTPRegionFunc(sum_exp_local)[0] : sum_exp_local; - // 4. Perform Softmax(local shards but normalize globally) + // 4. Perform Softmax (local shards but normalize globally). auto softmax_local = exp_local->Div(sum_exp->Unsqueeze(-1)); // 5. Perform allreduce to get global predicted_logit @@ -479,7 +556,7 @@ VocabParallelCrossEntropy::Forward(const std::vector> &i auto log_sum_exp = sum_exp->Log(); auto loss = log_sum_exp->Sub(predicted); - // 7. Label smoothing(According to Megatron-LM) + // 7. Apply label smoothing according to Megatron-LM. // TODO(zbl): adjust smoothing coef according to vocab_size_original if (label_smoothing_ > 0.0f) { // mean_logp over *valid tokens only*: diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 8a9b11cf4..e6fc1805b 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -4,6 +4,10 @@ #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/load_strategy.h" +#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/optimizer.h" @@ -14,8 +18,58 @@ using namespace infini_train; namespace nn = infini_train::nn; +namespace { + +class NamedParameterModule final : public nn::Module { +public: + void AddParameter(const std::string &name, const std::shared_ptr ¶meter) { + parameters_[name] = parameter; + } + + void AddModule(const std::string &name, const std::shared_ptr &module) { modules_[name] = module; } +}; + +} // namespace + +TEST(ModuleNamedParametersTest, SupportsTorchStyleArgumentsAndSharedParameterDeduplication) { + auto root = std::make_shared(); + auto child = std::make_shared(); + auto shared = std::make_shared(std::vector{1}, DataType::kFLOAT32, Device()); + auto child_weight = std::make_shared(std::vector{1}, DataType::kFLOAT32, Device()); + root->AddParameter("root_weight", shared); + child->AddParameter("alias", shared); + child->AddParameter("weight", child_weight); + root->AddModule("child", child); + + const auto local = root->NamedParameters("model", false); + ASSERT_EQ(local.size(), 1); + EXPECT_EQ(local[0].first, "model.root_weight"); + EXPECT_EQ(local[0].second, shared); + + const auto deduplicated = root->NamedParameters("model"); + ASSERT_EQ(deduplicated.size(), 2); + EXPECT_EQ(deduplicated[0].first, "model.root_weight"); + EXPECT_EQ(deduplicated[1].first, "model.child.weight"); + + const auto aliases = root->NamedParameters("model", true, false); + ASSERT_EQ(aliases.size(), 3); + EXPECT_EQ(aliases[0].first, "model.root_weight"); + EXPECT_EQ(aliases[1].first, "model.child.alias"); + EXPECT_EQ(aliases[1].second, shared); + EXPECT_EQ(aliases[2].first, "model.child.weight"); +} + class CheckpointSerializationTest : public test::InfiniTrainTest {}; +TEST(ShardedStateDictTest, RejectsDuplicateKeysWhenMerging) { + checkpoint::ShardedStateDict destination; + destination.tensors["weight"] = {.key = "weight"}; + checkpoint::ShardedStateDict source; + source.tensors["weight"] = {.key = "weight"}; + + EXPECT_DEATH(destination.Merge(std::move(source)), "Duplicate sharded state-dict key: weight"); +} + TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { auto dir = std::filesystem::temp_directory_path() / "test_ckpt_fp32"; std::filesystem::remove_all(dir); @@ -52,4 +106,350 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { std::filesystem::remove_all(dir); } +TEST_P(CheckpointSerializationTest, DirectMetadataOffsetSupportsColumnSlices) { + auto dir = std::filesystem::temp_directory_path() / "test_ckpt_region"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + auto matrix = std::make_shared(std::vector{4, 4}, DataType::kFLOAT32, Device()); + auto *values = static_cast(matrix->DataPtr()); + for (int row = 0; row < 4; ++row) { + for (int column = 0; column < 4; ++column) { values[row * 4 + column] = row * 10.0f + column; } + } + auto path = dir / "model.ckpt"; + Checkpoint::SaveStateDictFile(path, {{"matrix", matrix}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("matrix") - 1 + sizeof(int8_t) + + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + checkpoint::LoadPlan plan; + plan.tensors["matrix"] = {.key = "matrix", + .dtype = DataType::kFLOAT32, + .global_shape = {4, 4}, + .target_shape = {4, 2}, + .shard_dim = 1, + .reads = {{.key = "matrix", + .filename = "model.ckpt", + .dtype = DataType::kFLOAT32, + .global_shape = {4, 4}, + .byte_size = sizeof(float) * 16, + .data_offset = data_offset, + .shard_dim = 1, + .source_offset = 1, + .target_offset = 0, + .length = 2, + .source_shape = {4, 4}}}}; + checkpoint::IndexedRegionLoadStrategy strategy; + const auto planned = strategy.Execute(dir, plan); + const auto *planned_data = static_cast(planned.at("matrix")->DataPtr()); + for (int row = 0; row < 4; ++row) { + for (int column = 0; column < 2; ++column) { + EXPECT_FLOAT_EQ(planned_data[row * 2 + column], row * 10.0f + column + 1); + } + } + + std::filesystem::remove_all(dir); +} + +TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { + const auto dir = std::filesystem::temp_directory_path() / "test_vocab_padding_reshard"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + + auto source = std::make_shared(std::vector{5, 2}, DataType::kFLOAT32, Device()); + auto *source_data = static_cast(source->DataPtr()); + for (int i = 0; i < 10; ++i) { source_data[i] = static_cast(i); } + Checkpoint::SaveStateDictFile(dir / "model.ckpt", {{"lm_head.weight", source}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("lm_head.weight") - 1 + + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + + Checkpoint::CheckpointMetadata metadata; + metadata.tensors.push_back({.key = "lm_head.weight", + .dtype_str = "float32", + .global_shape = {5, 2}, + .local_shape = {5, 2}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .file = "model.ckpt", + .offset = data_offset, + .byte_size = sizeof(float) * 10}); + + checkpoint::ShardedStateDict target; + target.tensors["lm_head.weight"] = {.key = "lm_head.weight", + .dtype = DataType::kFLOAT32, + .global_shape = {8, 2}, + .local_shape = {4, 2}, + .global_offset = {4, 0}, + .axis_fragmentations = {2, 1}}; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + ASSERT_EQ(plan.tensors.at("lm_head.weight").trailing_zero_fill, 3); + checkpoint::IndexedRegionLoadStrategy strategy; + const auto loaded = strategy.Execute(dir, plan); + const auto *values = static_cast(loaded.at("lm_head.weight")->DataPtr()); + EXPECT_FLOAT_EQ(values[0], 8.0f); + EXPECT_FLOAT_EQ(values[1], 9.0f); + for (int i = 2; i < 8; ++i) { EXPECT_FLOAT_EQ(values[i], 0.0f); } + + std::filesystem::remove_all(dir); +} + +TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { + auto dir = std::filesystem::temp_directory_path() / "test_global_metadata"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + Checkpoint::CheckpointMetadata metadata; + metadata.version = 3; + metadata.iteration = 17; + metadata.has_metadata = true; + metadata.parallel_config = {.tp_size = 2, .pp_size = 2, .dp_size = 1, .sp_size = 1}; + metadata.tensors.push_back({.key = "layer.0.weight", + .dtype_str = "float32", + .global_shape = {8, 4}, + .local_shape = {4, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = {{.global_offset = 0, .local_offset = 0, .length = 4}}, + .file = "rank_000000/model.ckpt", + .byte_size = 64, + .stored_on_ranks = {0}, + .pp_rank = 0}); + Checkpoint::SaveMetadataFile(dir / "metadata.json", metadata); + + auto loaded = Checkpoint::LoadMetadata(dir); + ASSERT_TRUE(loaded.has_metadata); + EXPECT_EQ(loaded.iteration, 17); + EXPECT_EQ(loaded.parallel_config.tp_size, 2); + EXPECT_EQ(loaded.parallel_config.pp_size, 2); + ASSERT_EQ(loaded.tensors.size(), 1); + EXPECT_EQ(loaded.tensors[0].file, "rank_000000/model.ckpt"); + EXPECT_EQ(loaded.tensors[0].global_offset, std::vector({0, 0})); + EXPECT_EQ(loaded.tensors[0].axis_fragmentations, std::vector({2, 1})); + ASSERT_EQ(loaded.tensors[0].segments.size(), 1); + EXPECT_EQ(loaded.tensors[0].segments[0], + (checkpoint::ShardSegment{.global_offset = 0, .local_offset = 0, .length = 4})); + std::filesystem::remove_all(dir); +} + INFINI_TRAIN_REGISTER_TEST(CheckpointSerializationTest); + +namespace { +Checkpoint::CheckpointMetadata::TensorEntry MakeSavedShard(const std::string &key, int count, int index, + int64_t global_size, const std::string &file) { + return {.key = key, + .dtype_str = "float32", + .global_shape = {global_size, 4}, + .local_shape = {global_size / count, 4}, + .global_offset = {global_size / count * index, 0}, + .axis_fragmentations = {count, 1}, + .file = file}; +} + +checkpoint::ShardedStateDict MakeTarget(const std::string &key, int count, int index, int64_t global_size) { + checkpoint::ShardedStateDict target; + target.tensors[key] = {.key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {global_size, 4}, + .local_shape = {global_size / count, 4}, + .global_offset = {global_size / count * index, 0}, + .axis_fragmentations = {count, 1}}; + return target; +} +} // namespace + +TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { + checkpoint::ShardedStateDict model; + model.tensors["c_attn.weight"] = { + .key = "c_attn.weight", + .dtype = DataType::kFLOAT32, + .global_shape = {24, 4}, + .local_shape = {6, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {4, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 4}, + {.global_offset = 16, .local_offset = 4, .length = 1}, + {.global_offset = 20, .local_offset = 5, .length = 1}, + }, + }; + + auto moment = std::make_shared(std::vector{6, 4}, DataType::kFLOAT32, Device()); + auto step = std::make_shared(std::vector{}, DataType::kINT64, Device()); + std::unordered_map> optimizer_state = { + {"adam.m.c_attn.weight", moment}, + {"adam.v.c_attn.weight", moment}, + {"adam.t", step}, + }; + + const auto optimizer = checkpoint::BuildOptimizerShardedStateDict(model, optimizer_state); + ASSERT_EQ(optimizer.tensors.size(), 3); + const auto &m = optimizer.tensors.at("adam.m.c_attn.weight"); + EXPECT_EQ(m.global_shape, model.tensors.at("c_attn.weight").global_shape); + EXPECT_EQ(m.local_shape, model.tensors.at("c_attn.weight").local_shape); + EXPECT_EQ(m.segments, model.tensors.at("c_attn.weight").segments); + EXPECT_EQ(m.local_key, "adam.m.c_attn.weight"); + const auto &t = optimizer.tensors.at("adam.t"); + EXPECT_TRUE(t.global_shape.empty()); + EXPECT_TRUE(t.local_shape.empty()); +} + +TEST(CheckpointLoadPlannerTest, RejectsUnknownCheckpointDtype) { + auto metadata = Checkpoint::CheckpointMetadata{}; + metadata.tensors = {MakeSavedShard("weight", 1, 0, 16, "rank_0/model.ckpt")}; + metadata.tensors.front().dtype_str = "unknown_dtype"; + + EXPECT_DEATH(checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 1, 0, 16)), + "Unsupported checkpoint tensor dtype: unknown_dtype"); +} + +TEST(CheckpointLoadPlannerTest, TensorParallelTwoToFourReadsOnlyOverlap) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {MakeSavedShard("weight", 2, 0, 16, "rank_0/model.ckpt"), + MakeSavedShard("weight", 2, 1, 16, "rank_1/model.ckpt")}; + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 4, 1, 16)); + const auto &reads = plan.tensors.at("weight").reads; + ASSERT_EQ(reads.size(), 1); + EXPECT_EQ(reads[0].filename, "rank_0/model.ckpt"); + EXPECT_EQ(reads[0].source_offset, 4); + EXPECT_EQ(reads[0].target_offset, 0); + EXPECT_EQ(reads[0].length, 4); +} + +TEST(CheckpointLoadPlannerTest, TensorParallelFourToTwoReadsTwoOverlaps) { + Checkpoint::CheckpointMetadata metadata; + for (int index = 0; index < 4; ++index) { + metadata.tensors.push_back( + MakeSavedShard("weight", 4, index, 16, "rank_" + std::to_string(index) + "/model.ckpt")); + } + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 2, 1, 16)); + const auto &reads = plan.tensors.at("weight").reads; + ASSERT_EQ(reads.size(), 2); + EXPECT_EQ(reads[0].filename, "rank_2/model.ckpt"); + EXPECT_EQ(reads[0].target_offset, 0); + EXPECT_EQ(reads[1].filename, "rank_3/model.ckpt"); + EXPECT_EQ(reads[1].target_offset, 4); +} + +TEST(CheckpointLoadPlannerTest, UsesExplicitGlobalOffsetsForUnevenShards) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {{.key = "weight", + .dtype_str = "float32", + .global_shape = {8, 4}, + .local_shape = {3, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .file = "rank_0/model.ckpt"}, + {.key = "weight", + .dtype_str = "float32", + .global_shape = {8, 4}, + .local_shape = {5, 4}, + .global_offset = {3, 0}, + .axis_fragmentations = {2, 1}, + .file = "rank_1/model.ckpt"}}; + checkpoint::ShardedStateDict target; + target.tensors["weight"] = {.key = "weight", + .dtype = DataType::kFLOAT32, + .global_shape = {8, 4}, + .local_shape = {4, 4}, + .global_offset = {2, 0}, + .axis_fragmentations = {2, 1}}; + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto &reads = plan.tensors.at("weight").reads; + ASSERT_EQ(reads.size(), 2); + EXPECT_EQ(reads[0].source_offset, 2); + EXPECT_EQ(reads[0].target_offset, 0); + EXPECT_EQ(reads[0].length, 1); + EXPECT_EQ(reads[1].source_offset, 0); + EXPECT_EQ(reads[1].target_offset, 1); + EXPECT_EQ(reads[1].length, 3); +} + +TEST(CheckpointLoadPlannerTest, QkvSegmentsUseDimZeroWhenTpIsOne) { + Checkpoint::CheckpointMetadata metadata; + auto saved = MakeSavedShard("c_attn.weight", 1, 0, 12, "old_pp/model.ckpt"); + saved.local_shape = {12, 4}; + saved.axis_fragmentations = {1, 1}; + saved.segments = { + {.global_offset = 0, .local_offset = 0, .length = 8}, + {.global_offset = 8, .local_offset = 8, .length = 2}, + {.global_offset = 10, .local_offset = 10, .length = 2}, + }; + metadata.tensors.push_back(std::move(saved)); + + checkpoint::ShardedStateDict target; + target.tensors["c_attn.weight"] = { + .key = "c_attn.weight", + .dtype = DataType::kFLOAT32, + .global_shape = {12, 4}, + .local_shape = {12, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 8}, + {.global_offset = 8, .local_offset = 8, .length = 2}, + {.global_offset = 10, .local_offset = 10, .length = 2}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto &tensor_plan = plan.tensors.at("c_attn.weight"); + EXPECT_EQ(tensor_plan.shard_dim, 0); + ASSERT_EQ(tensor_plan.reads.size(), 3); + EXPECT_EQ(tensor_plan.reads[0].target_offset, 0); + EXPECT_EQ(tensor_plan.reads[1].target_offset, 8); + EXPECT_EQ(tensor_plan.reads[2].target_offset, 10); +} + +TEST(CheckpointLoadPlannerTest, QkvSegmentsPreserveTargetLocalLayoutAcrossTpChange) { + Checkpoint::CheckpointMetadata metadata; + for (int rank = 0; rank < 4; ++rank) { + auto shard = MakeSavedShard("c_attn.weight", 4, rank, 24, "rank_" + std::to_string(rank) + "/model.ckpt"); + shard.local_shape = {6, 4}; + shard.global_offset = {0, 0}; + shard.segments = { + {.global_offset = rank * 4, .local_offset = 0, .length = 4}, + {.global_offset = 16 + rank, .local_offset = 4, .length = 1}, + {.global_offset = 20 + rank, .local_offset = 5, .length = 1}, + }; + metadata.tensors.push_back(std::move(shard)); + } + + checkpoint::ShardedStateDict target; + target.tensors["c_attn.weight"] = { + .key = "c_attn.weight", + .dtype = DataType::kFLOAT32, + .global_shape = {24, 4}, + .local_shape = {12, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 8}, + {.global_offset = 16, .local_offset = 8, .length = 2}, + {.global_offset = 20, .local_offset = 10, .length = 2}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto &reads = plan.tensors.at("c_attn.weight").reads; + ASSERT_EQ(reads.size(), 6); + const std::vector files = {"rank_0/model.ckpt", "rank_1/model.ckpt", "rank_0/model.ckpt", + "rank_1/model.ckpt", "rank_0/model.ckpt", "rank_1/model.ckpt"}; + const std::vector source_offsets = {0, 0, 4, 4, 5, 5}; + const std::vector target_offsets = {0, 4, 8, 9, 10, 11}; + for (size_t i = 0; i < reads.size(); ++i) { + EXPECT_EQ(reads[i].filename, files[i]); + EXPECT_EQ(reads[i].source_offset, source_offsets[i]); + EXPECT_EQ(reads[i].target_offset, target_offsets[i]); + } +} + +TEST(CheckpointLoadPlannerTest, PipelineReshardPlansOnlyTargetStageKeys) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {MakeSavedShard("layer.0.weight", 1, 0, 8, "old_pp0/model.ckpt"), + MakeSavedShard("layer.1.weight", 1, 0, 8, "old_pp1/model.ckpt")}; + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("layer.1.weight", 1, 0, 8)); + ASSERT_EQ(plan.tensors.size(), 1); + ASSERT_EQ(plan.tensors.at("layer.1.weight").reads.size(), 1); + EXPECT_EQ(plan.tensors.at("layer.1.weight").reads[0].filename, "old_pp1/model.ckpt"); +} diff --git a/tests/checkpoint/test_optimizer_state.cc b/tests/checkpoint/test_optimizer_state.cc index 1cbb8b9f3..e3b0ebdfa 100644 --- a/tests/checkpoint/test_optimizer_state.cc +++ b/tests/checkpoint/test_optimizer_state.cc @@ -78,6 +78,25 @@ TEST_P(OptimizerStateTest, AdamStateDictRoundTrip) { } } +TEST_P(OptimizerStateTest, AdamStateDictUsesStableParameterNames) { + auto first = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); + auto second = std::make_shared(std::vector{3}, DataType::kFLOAT32, GetDevice()); + auto adam = std::make_shared(std::vector>{first, second}, 0.001); + adam->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"}); + + const auto state = adam->StateDict(); + EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight")); + EXPECT_TRUE(state.contains("adam.v.transformer.h.0.weight")); + EXPECT_TRUE(state.contains("adam.m.transformer.h.0.bias")); + EXPECT_TRUE(state.contains("adam.v.transformer.h.0.bias")); + EXPECT_TRUE(state.contains("adam.t")); + + auto restored = std::make_shared(std::vector>{first, second}, 0.001); + restored->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"}); + restored->LoadStateDict(state); + EXPECT_EQ(restored->StateDict().size(), state.size()); +} + // ---------- SGD ---------- TEST_P(OptimizerStateTest, SGDStateDictEmpty) { auto param = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); diff --git a/tests/transformer/test_transformer_architecture.cc b/tests/transformer/test_transformer_architecture.cc index 4cec471de..c754ea725 100644 --- a/tests/transformer/test_transformer_architecture.cc +++ b/tests/transformer/test_transformer_architecture.cc @@ -163,6 +163,10 @@ TEST_P(TransformerModuleTest, LLaMA3Model) { auto model = std::make_shared(config); model->To(GetDevice()); EXPECT_FALSE(model->Parameters().empty()); + const auto sharded_state = model->ShardedStateDict(); + for (const auto &[name, parameter] : model->NamedParameters()) { + EXPECT_TRUE(sharded_state.tensors.contains(name)) << "Missing shard metadata for named parameter: " << name; + } } TEST_P(TransformerModuleTest, RoPEUtils) { From 9ddcb1c590a13bad5d70b6c2e62ab59a97d9be73 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Wed, 12 Aug 2026 07:32:29 +0000 Subject: [PATCH 2/5] fix: complete checkpoint metadata for dtype and LoRA shards --- example/gpt2/main.cc | 1 + example/llama3/main.cc | 1 + infini_train/include/checkpoint/checkpoint.h | 2 + .../include/checkpoint/checkpoint_manager.h | 1 + .../include/nn/lora/lora_parallel_linear.h | 4 ++ infini_train/src/checkpoint/checkpoint.cc | 14 +++-- .../src/checkpoint/checkpoint_manager.cc | 8 ++- infini_train/src/checkpoint/load_strategy.cc | 6 +- infini_train/src/checkpoint/save_planner.cc | 1 + .../src/nn/lora/lora_parallel_linear.cc | 52 ++++++++++++++++ .../transformer/causal_self_attention.cc | 2 + .../src/nn/modules/transformer/transformer.cc | 24 +++++++- .../test_checkpoint_serialization.cc | 59 ++++++++++++++++++- tests/checkpoint/test_trainer_state.cc | 3 + tests/lora/test_lora.cc | 29 +++++++++ 15 files changed, 197 insertions(+), 10 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 7697ad213..ec41c7465 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -427,6 +427,7 @@ void Train(const nn::parallel::Rank &rank) { .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, diff --git a/example/llama3/main.cc b/example/llama3/main.cc index a9f405fa5..12fecce2b 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -417,6 +417,7 @@ void Train(const nn::parallel::Rank &rank) { .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index 586c1e26c..d93ea35f3 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -32,6 +32,7 @@ struct TrainerState { int tp_size = 1; int sp_size = 1; int pp_size = 1; + int vpp_size = 1; }; class Checkpoint { @@ -63,6 +64,7 @@ class Checkpoint { int pp_size = 1; int dp_size = 1; int sp_size = 1; + int vpp_size = 1; } parallel_config; struct TensorEntry { diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 47e07fdd6..d90bf655f 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -49,6 +49,7 @@ struct SaveCheckpointArgs { int tp_size = 1; int sp_size = 1; int pp_size = 1; + int vpp_size = 1; std::filesystem::path checkpoint_root_dir; size_t max_checkpoint_keep = 0; const nn::parallel::Rank &rank; diff --git a/infini_train/include/nn/lora/lora_parallel_linear.h b/infini_train/include/nn/lora/lora_parallel_linear.h index d73a6e2b5..b485e3c7c 100644 --- a/infini_train/include/nn/lora/lora_parallel_linear.h +++ b/infini_train/include/nn/lora/lora_parallel_linear.h @@ -34,6 +34,8 @@ class LoRAColumnParallelLinear : public nn::parallel::ColumnParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + void MergeWeights(); void UnmergeWeights(); bool IsMerged() const; @@ -74,6 +76,8 @@ class LoRARowParallelLinear : public nn::parallel::RowParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + void MergeWeights(); void UnmergeWeights(); bool IsMerged() const; diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index ef3559fb4..b76f5b410 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -241,7 +241,8 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m LOG(ERROR) << "[CKPT] Load done: global_step=" << state.global_step << ", consumed_train_samples=" << state.consumed_train_samples << ", topology(ddp,tp,sp,pp)=(" - << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; + << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << "," + << state.vpp_size << ")"; } Checkpoint::SavedTensorLocations @@ -335,7 +336,8 @@ void Checkpoint::SaveTrainerState(const std::filesystem::path &path, const Train ofs << " \"ddp_size\": " << state.ddp_size << ",\n"; ofs << " \"tp_size\": " << state.tp_size << ",\n"; ofs << " \"sp_size\": " << state.sp_size << ",\n"; - ofs << " \"pp_size\": " << state.pp_size << "\n"; + ofs << " \"pp_size\": " << state.pp_size << ",\n"; + ofs << " \"vpp_size\": " << state.vpp_size << "\n"; ofs << "}\n"; } @@ -357,6 +359,7 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.tp_size = ExtractNumberField(content, "tp_size", 1); state.sp_size = ExtractNumberField(content, "sp_size", 1); state.pp_size = ExtractNumberField(content, "pp_size", 1); + state.vpp_size = ExtractNumberField(content, "vpp_size", 1); return state; } @@ -447,7 +450,8 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, ofs << " \"tp_size\": " << state.tp_size << ",\n"; ofs << " \"pp_size\": " << state.pp_size << ",\n"; ofs << " \"dp_size\": " << state.ddp_size << ",\n"; - ofs << " \"sp_size\": " << state.sp_size << "\n"; + ofs << " \"sp_size\": " << state.sp_size << ",\n"; + ofs << " \"vpp_size\": " << state.vpp_size << "\n"; ofs << " },\n"; ofs << " \"model_config\": {\n"; ofs << " \"n_layer\": " << state.n_layer << ",\n"; @@ -569,6 +573,7 @@ static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem:: meta.parallel_config.pp_size = ExtractNumberField(content, "pp_size", 1); meta.parallel_config.dp_size = ExtractNumberField(content, "dp_size", 1); meta.parallel_config.sp_size = ExtractNumberField(content, "sp_size", 1); + meta.parallel_config.vpp_size = ExtractNumberField(content, "vpp_size", 1); // Locate the tensors array. auto tensors_key = content.find("\"tensors\""); @@ -771,7 +776,8 @@ void Checkpoint::SaveMetadataFile(const std::filesystem::path &path, const Check ofs << " \"tp_size\": " << metadata.parallel_config.tp_size << ",\n"; ofs << " \"pp_size\": " << metadata.parallel_config.pp_size << ",\n"; ofs << " \"dp_size\": " << metadata.parallel_config.dp_size << ",\n"; - ofs << " \"sp_size\": " << metadata.parallel_config.sp_size << "\n"; + ofs << " \"sp_size\": " << metadata.parallel_config.sp_size << ",\n"; + ofs << " \"vpp_size\": " << metadata.parallel_config.vpp_size << "\n"; ofs << " },\n"; ofs << " \"tensors\": [\n"; for (size_t i = 0; i < metadata.tensors.size(); ++i) { diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index 6bc771473..9db2c1514 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -17,6 +17,7 @@ #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" +#include "infini_train/include/nn/parallel/ddp/distributed_optimizer.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/nn/parallel/parallel_functional.h" #include "infini_train/include/nn/parallel/work.h" @@ -94,6 +95,8 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & LOG(INFO) << "No checkpoint specified for resume. Starting training from scratch."; return result; } + CHECK(dynamic_cast(args.optimizer.get()) == nullptr) + << "Checkpoint restore does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); CHECK(std::filesystem::exists(checkpoint_dir / "metadata.json")) @@ -121,6 +124,8 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & } void SaveCheckpoint(const SaveCheckpointArgs &args) { + CHECK(dynamic_cast(args.optimizer) == nullptr) + << "Checkpoint save does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; const auto checkpoint_start = std::chrono::high_resolution_clock::now(); TrainerState state{.global_step = args.global_step, .consumed_train_samples = static_cast(args.consumed_train_samples), @@ -132,7 +137,8 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { .ddp_size = args.ddp_size, .tp_size = args.tp_size, .sp_size = args.sp_size, - .pp_size = args.pp_size}; + .pp_size = args.pp_size, + .vpp_size = args.vpp_size}; const auto iteration_dir = args.checkpoint_root_dir.empty() ? args.save_dir : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); diff --git a/infini_train/src/checkpoint/load_strategy.cc b/infini_train/src/checkpoint/load_strategy.cc index 476637881..35e967c3b 100644 --- a/infini_train/src/checkpoint/load_strategy.cc +++ b/infini_train/src/checkpoint/load_strategy.cc @@ -100,7 +100,11 @@ LoadedStateDict IndexedRegionLoadStrategy::Execute(const std::filesystem::path & CHECK_GT(read.data_offset, 0) << "Checkpoint metadata lacks a valid tensor data offset for " << key << "; regenerate the checkpoint with the current format"; auto &stream = GetFile(file_cache, checkpoint_dir, read.filename); - pieces.push_back(read.shard_dim < 0 ? ReadTensor(stream, read) : ReadTensorRegion(stream, read)); + auto piece = read.shard_dim < 0 ? ReadTensor(stream, read) : ReadTensorRegion(stream, read); + if (piece->Dtype() != tensor_plan.dtype) { + piece = std::make_shared(piece->To(tensor_plan.dtype)); + } + pieces.push_back(std::move(piece)); } if (tensor_plan.trailing_zero_fill > 0) { diff --git a/infini_train/src/checkpoint/save_planner.cc b/infini_train/src/checkpoint/save_planner.cc index 17ae1a85b..48e61c4ac 100644 --- a/infini_train/src/checkpoint/save_planner.cc +++ b/infini_train/src/checkpoint/save_planner.cc @@ -40,6 +40,7 @@ BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, auto info = model_it->second; info.key = key; info.local_key = key; + info.dtype = tensor->Dtype(); result.tensors.emplace(key, std::move(info)); } return result; diff --git a/infini_train/src/nn/lora/lora_parallel_linear.cc b/infini_train/src/nn/lora/lora_parallel_linear.cc index 9b038e2d4..298b3e074 100644 --- a/infini_train/src/nn/lora/lora_parallel_linear.cc +++ b/infini_train/src/nn/lora/lora_parallel_linear.cc @@ -219,6 +219,32 @@ std::vector> LoRAColumnParallelLinear::LoRAParameters() return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } +checkpoint::ShardedStateDict LoRAColumnParallelLinear::ShardedStateDict(const std::string &prefix) const { + auto state = parallel::ColumnParallelLinear::ShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + + const auto &lora_a = parameter(kParamLoraAName); + checkpoint::ShardedTensor a; + a.key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; + a.dtype = lora_a->Dtype(); + a.global_shape = lora_a->Dims(); + a.local_shape = lora_a->Dims(); + a.global_offset = {0, 0}; + a.axis_fragmentations = {1, 1}; + state.tensors.emplace(a.key, std::move(a)); + + const auto &lora_b = parameter(kParamLoraBName); + checkpoint::ShardedTensor b; + b.key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; + b.dtype = lora_b->Dtype(); + b.global_shape = {lora_b->Dims()[0] * tp_size, lora_b->Dims()[1]}; + b.local_shape = lora_b->Dims(); + b.global_offset = {lora_b->Dims()[0] * parallel::tp_rank, 0}; + b.axis_fragmentations = {tp_size, 1}; + state.tensors.emplace(b.key, std::move(b)); + return state; +} + bool LoRAColumnParallelLinear::IsMerged() const { return merged_; } int64_t LoRAColumnParallelLinear::in_features() const { return in_features_; } @@ -429,6 +455,32 @@ std::vector> LoRARowParallelLinear::LoRAParameters() con return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } +checkpoint::ShardedStateDict LoRARowParallelLinear::ShardedStateDict(const std::string &prefix) const { + auto state = parallel::RowParallelLinear::ShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + + const auto &lora_a = parameter(kParamLoraAName); + checkpoint::ShardedTensor a; + a.key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; + a.dtype = lora_a->Dtype(); + a.global_shape = {lora_a->Dims()[0], lora_a->Dims()[1] * tp_size}; + a.local_shape = lora_a->Dims(); + a.global_offset = {0, lora_a->Dims()[1] * parallel::tp_rank}; + a.axis_fragmentations = {1, tp_size}; + state.tensors.emplace(a.key, std::move(a)); + + const auto &lora_b = parameter(kParamLoraBName); + checkpoint::ShardedTensor b; + b.key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; + b.dtype = lora_b->Dtype(); + b.global_shape = lora_b->Dims(); + b.local_shape = lora_b->Dims(); + b.global_offset = {0, 0}; + b.axis_fragmentations = {1, 1}; + state.tensors.emplace(b.key, std::move(b)); + return state; +} + bool LoRARowParallelLinear::IsMerged() const { return merged_; } int64_t LoRARowParallelLinear::in_features() const { return in_features_; } diff --git a/infini_train/src/nn/modules/transformer/causal_self_attention.cc b/infini_train/src/nn/modules/transformer/causal_self_attention.cc index 5db2ed1d9..57a341fbc 100644 --- a/infini_train/src/nn/modules/transformer/causal_self_attention.cc +++ b/infini_train/src/nn/modules/transformer/causal_self_attention.cc @@ -10,6 +10,7 @@ #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/init.h" +#include "infini_train/include/nn/lora/lora_parallel_linear.h" #include "infini_train/include/nn/modules/normalization.h" #include "infini_train/include/nn/modules/sparse.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" @@ -112,6 +113,7 @@ checkpoint::ShardedStateDict CausalSelfAttention::ShardedStateDict(const std::st if (config_.add_bias_linear) { set_qkv_segments(parallel::ColumnParallelLinear::kParamBiasName); } + set_qkv_segments(lora::LoRAColumnParallelLinear::kParamLoraBName); return state; } diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index d6e6e28a0..9704f9741 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -329,12 +329,30 @@ checkpoint::ShardedStateDict TransformerModel::ShardedStateDict(const std::strin std::vector>> TransformerModel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { - auto parameters = Module::NamedParameters(prefix, recurse, remove_duplicate); + if (!recurse) { + return Module::NamedParameters(prefix, false, remove_duplicate); + } + + // Select public aliases so optimizer state keys match ShardedStateDict keys. + auto parameters = Module::NamedParameters(prefix, true, false); + const auto sharded_state = ShardedStateDict(prefix); const auto global_layers = GlobalLayerIndices(stage_info_); std::vector local_layers(global_layers.size()); std::iota(local_layers.begin(), local_layers.end(), 0); - for (auto &[name, parameter] : parameters) { name = RemapLayerKey(name, local_layers, global_layers); } - return parameters; + + std::vector>> result; + std::unordered_set visited; + for (auto &[name, parameter] : parameters) { + name = RemapLayerKey(name, local_layers, global_layers); + if (!sharded_state.tensors.contains(name)) { + continue; + } + if (remove_duplicate && !visited.insert(parameter.get()).second) { + continue; + } + result.emplace_back(std::move(name), std::move(parameter)); + } + return result; } void TransformerModel::LoadStateDict(const std::unordered_map> &state_dict) { diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index e6fc1805b..000e80e71 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -148,6 +148,45 @@ TEST_P(CheckpointSerializationTest, DirectMetadataOffsetSupportsColumnSlices) { std::filesystem::remove_all(dir); } +TEST_P(CheckpointSerializationTest, ConvertsSavedBF16TensorToFP32Target) { + auto dir = std::filesystem::temp_directory_path() / "test_ckpt_bf16_to_fp32"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + + auto source_fp32 = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, Device()); + auto *source_data = static_cast(source_fp32->DataPtr()); + source_data[0] = 1.0f; + source_data[1] = 2.0f; + source_data[2] = 3.0f; + source_data[3] = 4.0f; + auto source_bf16 = std::make_shared(source_fp32->To(DataType::kBFLOAT16)); + Checkpoint::SaveStateDictFile(dir / "model.ckpt", {{"weight", source_bf16}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("weight") - 1 + + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + + checkpoint::LoadPlan plan; + plan.tensors["weight"] = {.key = "weight", + .dtype = DataType::kFLOAT32, + .global_shape = {2, 2}, + .target_shape = {2, 2}, + .reads = {{.key = "weight", + .filename = "model.ckpt", + .dtype = DataType::kBFLOAT16, + .global_shape = {2, 2}, + .byte_size = source_bf16->SizeInBytes(), + .data_offset = data_offset, + .shard_dim = -1, + .source_shape = {2, 2}}}}; + + checkpoint::IndexedRegionLoadStrategy strategy; + const auto loaded = strategy.Execute(dir, plan).at("weight"); + ASSERT_EQ(loaded->Dtype(), DataType::kFLOAT32); + const auto *loaded_data = static_cast(loaded->DataPtr()); + for (int i = 0; i < 4; ++i) { EXPECT_FLOAT_EQ(loaded_data[i], source_data[i]); } + + std::filesystem::remove_all(dir); +} + TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { const auto dir = std::filesystem::temp_directory_path() / "test_vocab_padding_reshard"; std::filesystem::remove_all(dir); @@ -199,7 +238,7 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { metadata.version = 3; metadata.iteration = 17; metadata.has_metadata = true; - metadata.parallel_config = {.tp_size = 2, .pp_size = 2, .dp_size = 1, .sp_size = 1}; + metadata.parallel_config = {.tp_size = 2, .pp_size = 2, .dp_size = 1, .sp_size = 1, .vpp_size = 2}; metadata.tensors.push_back({.key = "layer.0.weight", .dtype_str = "float32", .global_shape = {8, 4}, @@ -218,6 +257,7 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { EXPECT_EQ(loaded.iteration, 17); EXPECT_EQ(loaded.parallel_config.tp_size, 2); EXPECT_EQ(loaded.parallel_config.pp_size, 2); + EXPECT_EQ(loaded.parallel_config.vpp_size, 2); ASSERT_EQ(loaded.tensors.size(), 1); EXPECT_EQ(loaded.tensors[0].file, "rank_000000/model.ckpt"); EXPECT_EQ(loaded.tensors[0].global_offset, std::vector({0, 0})); @@ -285,11 +325,28 @@ TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { EXPECT_EQ(m.local_shape, model.tensors.at("c_attn.weight").local_shape); EXPECT_EQ(m.segments, model.tensors.at("c_attn.weight").segments); EXPECT_EQ(m.local_key, "adam.m.c_attn.weight"); + EXPECT_EQ(m.dtype, moment->Dtype()); const auto &t = optimizer.tensors.at("adam.t"); EXPECT_TRUE(t.global_shape.empty()); EXPECT_TRUE(t.local_shape.empty()); } +TEST(CheckpointOptimizerShardingTest, AdamMomentUsesOptimizerStateDtype) { + checkpoint::ShardedStateDict model; + model.tensors["weight"] = {.key = "weight", + .dtype = DataType::kBFLOAT16, + .global_shape = {4, 4}, + .local_shape = {4, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}}; + auto moment = std::make_shared(std::vector{4, 4}, DataType::kFLOAT32, Device()); + + const auto optimizer + = checkpoint::BuildOptimizerShardedStateDict(model, {{"adam.m.weight", moment}, {"adam.v.weight", moment}}); + EXPECT_EQ(optimizer.tensors.at("adam.m.weight").dtype, DataType::kFLOAT32); + EXPECT_EQ(optimizer.tensors.at("adam.v.weight").dtype, DataType::kFLOAT32); +} + TEST(CheckpointLoadPlannerTest, RejectsUnknownCheckpointDtype) { auto metadata = Checkpoint::CheckpointMetadata{}; metadata.tensors = {MakeSavedShard("weight", 1, 0, 16, "rank_0/model.ckpt")}; diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index ec4d61e8b..efcaf8b3f 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -31,6 +31,7 @@ TEST_P(TrainerStateTest, DefaultValues) { EXPECT_EQ(state.tp_size, 1); EXPECT_EQ(state.sp_size, 1); EXPECT_EQ(state.pp_size, 1); + EXPECT_EQ(state.vpp_size, 1); } TEST_P(TrainerStateTest, TrainerStateFileCreated) { @@ -73,6 +74,7 @@ TEST_P(TrainerStateTest, RoundTrip) { .tp_size = 1, .sp_size = 1, .pp_size = 2, + .vpp_size = 4, }; auto model1 = std::make_shared(1, 3, true, GetDevice()); @@ -101,6 +103,7 @@ TEST_P(TrainerStateTest, RoundTrip) { EXPECT_EQ(loaded.vocab_size, 128256); EXPECT_EQ(loaded.ddp_size, 2); EXPECT_EQ(loaded.pp_size, 2); + EXPECT_EQ(loaded.vpp_size, 4); std::filesystem::remove_all(dir); } diff --git a/tests/lora/test_lora.cc b/tests/lora/test_lora.cc index 26cffdcaa..70d83a583 100644 --- a/tests/lora/test_lora.cc +++ b/tests/lora/test_lora.cc @@ -7,11 +7,13 @@ #include "infini_train/include/nn/lora/lora_config.h" #include "infini_train/include/nn/lora/lora_linear.h" +#include "infini_train/include/nn/lora/lora_parallel_linear.h" #include "infini_train/include/nn/lora/lora_utils.h" #include "infini_train/include/nn/modules/container.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/tensor_parallel.h" #include "infini_train/include/tensor.h" #include "tests/common/test_utils.h" @@ -78,6 +80,33 @@ TEST_P(LoRATest, LoRAConfigScaling) { EXPECT_EQ(config.Scaling(), expected_scaling); } +TEST_P(LoRATest, ParallelLoRAShardedStateDictIncludesAdapterParameters) { + LoRAConfig config; + config.rank = 2; + + auto column_base = std::make_shared( + 4, 6, /*bias=*/false, /*gather_output=*/false, /*input_is_parallel=*/false, /*skip_bias_add=*/false, + /*sequence_parallel=*/false); + auto column = std::make_shared(column_base, config, 4, 6); + const auto column_state = column->ShardedStateDict("column"); + ASSERT_TRUE(column_state.tensors.contains("column.lora_A")); + ASSERT_TRUE(column_state.tensors.contains("column.lora_B")); + EXPECT_EQ(column_state.tensors.at("column.lora_A").axis_fragmentations, (std::vector{1, 1})); + EXPECT_EQ(column_state.tensors.at("column.lora_B").axis_fragmentations, + (std::vector{nn::parallel::global::GetTensorParallelSize(), 1})); + + auto row_base = std::make_shared( + 4, 6, /*bias=*/false, /*reduce_output=*/true, /*input_is_parallel=*/true, /*skip_bias_add=*/false, + /*sequence_parallel=*/false); + auto row = std::make_shared(row_base, config, 4, 6); + const auto row_state = row->ShardedStateDict("row"); + ASSERT_TRUE(row_state.tensors.contains("row.lora_A")); + ASSERT_TRUE(row_state.tensors.contains("row.lora_B")); + EXPECT_EQ(row_state.tensors.at("row.lora_A").axis_fragmentations, + (std::vector{1, nn::parallel::global::GetTensorParallelSize()})); + EXPECT_EQ(row_state.tensors.at("row.lora_B").axis_fragmentations, (std::vector{1, 1})); +} + TEST_P(LoRATest, PackedQKVShardGPTStyle) { auto full_qkv = MakeRowLabeledTensor(/*rows=*/12, /*cols=*/3, GetDevice()); auto shard = infini_train::nn::lora::detail::SlicePackedQKVRowsForTensorParallel(full_qkv, /*q_rows=*/4, From eb87cedbc83ca3613e24c754efa7e5b4c6215b16 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Fri, 14 Aug 2026 09:25:16 +0000 Subject: [PATCH 3/5] fix: update checkpoint tests for named optimizer API --- infini_train/include/nn/parallel/tensor_parallel.h | 2 ++ infini_train/src/checkpoint/checkpoint_manager.cc | 9 +++++++++ infini_train/src/checkpoint/reshard.cc | 5 +++++ infini_train/src/checkpoint/save_planner.cc | 2 +- infini_train/src/nn/parallel/tensor_parallel.cc | 3 ++- scripts/run_models_and_profile.bash | 7 ++++--- tests/checkpoint/test_optimizer_state.cc | 7 +++---- 7 files changed, 26 insertions(+), 9 deletions(-) diff --git a/infini_train/include/nn/parallel/tensor_parallel.h b/infini_train/include/nn/parallel/tensor_parallel.h index 06a0abea6..4fbedb2a5 100644 --- a/infini_train/include/nn/parallel/tensor_parallel.h +++ b/infini_train/include/nn/parallel/tensor_parallel.h @@ -95,6 +95,8 @@ class VocabParallelEmbedding : public nn::CloneableModule(args.optimizer.get()) == nullptr) << "Checkpoint restore does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; + // Resolve the checkpoint generation and load the global shard metadata. auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); CHECK(std::filesystem::exists(checkpoint_dir / "metadata.json")) << "Checkpoint metadata.json not found: " << checkpoint_dir; @@ -105,9 +106,11 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & CHECK(metadata.has_metadata); CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + // Reconstruct model and optimizer state for the current parallel topology. checkpoint::LoadDistributedCheckpoint(checkpoint_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get(), metadata); + // Validate architecture invariants before restoring training progress. CHECK_EQ(args.state.n_layer, args.model_config.n_layer); CHECK_EQ(args.state.n_head, args.model_config.n_head); CHECK_EQ(args.state.n_kv_head, args.model_config.n_kv_head); @@ -127,6 +130,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { CHECK(dynamic_cast(args.optimizer) == nullptr) << "Checkpoint save does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; const auto checkpoint_start = std::chrono::high_resolution_clock::now(); + // Snapshot training progress and the topology that produced this checkpoint. TrainerState state{.global_step = args.global_step, .consumed_train_samples = static_cast(args.consumed_train_samples), .n_layer = args.n_layer, @@ -144,6 +148,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); std::filesystem::create_directories(iteration_dir); + // Reset the manifest staging area before writer ranks publish their metadata. const auto staging_root = iteration_dir / ".metadata_tmp"; if (args.rank.IsMainRank()) { std::filesystem::remove_all(staging_root); @@ -152,12 +157,14 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { int dp_rank = 0, tp_rank = 0, pp_rank = 0; nn::parallel::global::GetCoordOf(args.rank.GlobalRank(), dp_rank, tp_rank, pp_rank); + // DP ranks hold replicas; only one DP replica writes each TP/PP shard. if (dp_rank != 0) { return; } const auto rank_dir = iteration_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); std::filesystem::create_directories(rank_dir); + // Describe logical shards, plan their physical layout, and write this rank shard. auto sharded_state = args.model.ShardedStateDict(); std::unordered_map> optimizer_state; if (args.optimizer != nullptr) { @@ -171,6 +178,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { const auto staging_rank_dir = staging_root / std::format("rank_{:06d}", args.rank.GlobalRank()); std::filesystem::create_directories(staging_rank_dir); + // Stage the local manifest until all writer ranks have completed their shards. const auto local_manifest = staging_rank_dir / "metadata.json"; if (std::filesystem::exists(local_manifest)) { std::filesystem::remove(local_manifest); @@ -182,6 +190,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { if (args.lr_scheduler != nullptr) { Checkpoint::SaveLRSchedulerStateFile(iteration_dir / "lr_scheduler.ckpt", args.lr_scheduler->StateDict()); } + // Aggregate writer manifests and atomically publish the global metadata. WaitForWriterManifests(staging_root, args.tp_size, args.pp_size, args.global_step); auto global_metadata = Checkpoint::LoadMetadata(staging_root); CHECK(global_metadata.has_metadata); diff --git a/infini_train/src/checkpoint/reshard.cc b/infini_train/src/checkpoint/reshard.cc index fb8842f8a..fb2b364ca 100644 --- a/infini_train/src/checkpoint/reshard.cc +++ b/infini_train/src/checkpoint/reshard.cc @@ -20,12 +20,15 @@ void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn:: const Checkpoint::CheckpointMetadata &metadata) { CHECK(metadata.has_metadata); CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + // Build this rank's target shard layout and plan overlap reads from the saved source shards. auto model_sharded_state = model.ShardedStateDict(); auto plan = LoadPlanner::PlanReshard(metadata, model_sharded_state); + // Execute the read plan, assemble target tensors, and load the reconstructed model state. IndexedRegionLoadStrategy strategy; auto result = strategy.Execute(checkpoint_dir, plan); model.LoadStateDict(result); + // Restore training progress, but rewrite topology fields to describe the current runtime. state = Checkpoint::LoadTrainerStateFile(checkpoint_dir / "trainer_state.json"); const int current_tp = nn::parallel::global::GetTensorParallelSize(); const int current_pp = nn::parallel::global::GetPipelineParallelSize(); @@ -36,6 +39,7 @@ void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn:: state.ddp_size = nn::parallel::global::GetDataParallelSize(); state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; + // Reshard optimizer tensors only when TP or PP changed; otherwise load the matching writer shard directly. if (optimizer != nullptr) { if (topology_changed) { const auto initialized_optimizer_state = optimizer->StateDict(); @@ -60,6 +64,7 @@ void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn:: optimizer->LoadStateDict(Checkpoint::LoadStateDictFile(optimizer_path)); } } + // Scheduler state is topology-independent and can be restored directly. if (lr_scheduler != nullptr && std::filesystem::exists(checkpoint_dir / "lr_scheduler.ckpt")) { lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(checkpoint_dir / "lr_scheduler.ckpt")); } diff --git a/infini_train/src/checkpoint/save_planner.cc b/infini_train/src/checkpoint/save_planner.cc index 48e61c4ac..8b61b0b17 100644 --- a/infini_train/src/checkpoint/save_planner.cc +++ b/infini_train/src/checkpoint/save_planner.cc @@ -36,7 +36,7 @@ BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, auto model_it = model_state.tensors.find(parameter_key); CHECK(model_it != model_state.tensors.end()) << "Optimizer state " << key << " has no matching named model parameter. " - << "Optimizer resharding requires set_parameter_names()."; + << "Optimizer resharding requires named parameters."; auto info = model_it->second; info.key = key; info.local_key = key; diff --git a/infini_train/src/nn/parallel/tensor_parallel.cc b/infini_train/src/nn/parallel/tensor_parallel.cc index b5a54b8d9..2755c15fa 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -401,7 +401,8 @@ checkpoint::ShardedStateDict RowParallelLinear::ShardedStateDict(const std::stri VocabParallelEmbedding::VocabParallelEmbedding(int64_t num_embeddings, int64_t embedding_dim, bool reduce_scatter_embeddings) - : CloneableModule(kType), embedding_dim_(embedding_dim), reduce_scatter_embeddings_(reduce_scatter_embeddings) { + : CloneableModule(kType), vocab_size_global_(num_embeddings), embedding_dim_(embedding_dim), + reduce_scatter_embeddings_(reduce_scatter_embeddings) { auto tp_size = global::GetTensorParallelSize(); CHECK_GT(tp_size, 0) << "No available devices found for VocabParallelEmbedding"; CHECK_GT(num_embeddings, 0); diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index 34edd26ff..84c6351f6 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -374,9 +374,10 @@ args_string_for_test() { jq -r --argjson g "$group_idx" --argjson t "$test_idx" --arg model "$model_name" --arg test_id "$test_id" ' def namespaced_path($p; $model; $mode): - if ($p | test("/checkpoint_step_[0-9]+($|/)")) then - ($p | capture("^(?.*)/(?checkpoint_step_[0-9]+(?:/.*)?)$")) as $m - | ($m.prefix + "/" + $model + "/" + $mode + "/" + $m.step) + if ($p | test("/(?:checkpoint_step_|iter_)[0-9]+($|/)")) then + ($p | capture("^(?.*)/(?:checkpoint_step_|iter_)(?[0-9]+)(?/.*)?$")) as $m + | ("0000000" + $m.iteration)[-7:] as $iteration + | ($m.prefix + "/" + $model + "/" + $mode + "/iter_" + $iteration + ($m.suffix // "")) else ($p + "/" + $model + "/" + $mode) end; diff --git a/tests/checkpoint/test_optimizer_state.cc b/tests/checkpoint/test_optimizer_state.cc index e3b0ebdfa..aa61b3354 100644 --- a/tests/checkpoint/test_optimizer_state.cc +++ b/tests/checkpoint/test_optimizer_state.cc @@ -81,8 +81,8 @@ TEST_P(OptimizerStateTest, AdamStateDictRoundTrip) { TEST_P(OptimizerStateTest, AdamStateDictUsesStableParameterNames) { auto first = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); auto second = std::make_shared(std::vector{3}, DataType::kFLOAT32, GetDevice()); - auto adam = std::make_shared(std::vector>{first, second}, 0.001); - adam->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"}); + const NamedParameterList named_parameters{{"transformer.h.0.weight", first}, {"transformer.h.0.bias", second}}; + auto adam = std::make_shared(named_parameters, 0.001); const auto state = adam->StateDict(); EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight")); @@ -91,8 +91,7 @@ TEST_P(OptimizerStateTest, AdamStateDictUsesStableParameterNames) { EXPECT_TRUE(state.contains("adam.v.transformer.h.0.bias")); EXPECT_TRUE(state.contains("adam.t")); - auto restored = std::make_shared(std::vector>{first, second}, 0.001); - restored->set_parameter_names({"transformer.h.0.weight", "transformer.h.0.bias"}); + auto restored = std::make_shared(named_parameters, 0.001); restored->LoadStateDict(state); EXPECT_EQ(restored->StateDict().size(), state.size()); } From 19ecce606970df0a8c9569356964936e7108e20d Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Fri, 9 Oct 2026 06:08:53 +0000 Subject: [PATCH 4/5] refactor: address checkpoint resharding review feedback --- example/gpt2/checkpoint_loader.cc | 66 ++--- example/gpt2/main.cc | 5 +- example/llama3/checkpoint_loader.cc | 32 +-- example/llama3/main.cc | 5 +- example/mixtral/main.cc | 2 + example/qwen3/checkpoint_loader.cc | 42 +--- example/qwen3/main.cc | 5 +- infini_train/include/checkpoint/checkpoint.h | 36 +-- .../include/checkpoint/checkpoint_manager.h | 3 +- infini_train/include/checkpoint/constants.h | 14 ++ .../include/checkpoint/load_planner.h | 2 +- infini_train/include/checkpoint/reshard.h | 22 -- .../include/checkpoint/save_planner.h | 38 +-- infini_train/include/core/ccl/ccl.h | 3 + .../include/nn/lora/lora_parallel_linear.h | 4 +- infini_train/include/nn/modules/module.h | 4 +- .../transformer/causal_self_attention.h | 2 +- .../nn/modules/transformer/transformer.h | 3 +- .../parallel/ddp/distributed_data_parallel.h | 4 +- .../nn/parallel/ddp/distributed_optimizer.h | 3 - .../nn/parallel/ddp/param_and_grad_buffer.h | 10 + .../include/nn/parallel/parallel_functional.h | 3 + .../nn/parallel/pp/pipeline_parallel.h | 2 +- .../include/nn/parallel/process_group.h | 16 ++ .../include/nn/parallel/tensor_parallel.h | 8 +- infini_train/include/nn/parallel/utils.h | 5 + infini_train/include/optimizer.h | 6 + .../include/{checkpoint => }/shard_spec.h | 18 +- infini_train/include/tensor.h | 4 + infini_train/src/checkpoint/checkpoint.cc | 233 +++++++++--------- .../src/checkpoint/checkpoint_manager.cc | 172 +++---------- infini_train/src/checkpoint/load_planner.cc | 37 ++- infini_train/src/checkpoint/reshard.cc | 76 ------ infini_train/src/checkpoint/save_planner.cc | 31 +-- infini_train/src/core/ccl/ccl.cc | 5 + infini_train/src/core/ccl/cuda/nccl_impl.cc | 42 ++++ infini_train/src/core/ccl/cuda/nccl_impl.h | 3 + .../src/nn/lora/lora_parallel_linear.cc | 30 +-- infini_train/src/nn/modules/module.cc | 26 +- infini_train/src/nn/modules/normalization.cc | 8 + .../transformer/causal_self_attention.cc | 10 +- .../src/nn/modules/transformer/moe/router.cc | 7 + .../src/nn/modules/transformer/transformer.cc | 95 ++----- .../parallel/ddp/distributed_data_parallel.cc | 8 +- .../nn/parallel/ddp/distributed_optimizer.cc | 18 +- .../nn/parallel/ddp/param_and_grad_buffer.cc | 6 + .../src/nn/parallel/parallel_functional.cc | 9 + .../src/nn/parallel/pp/pipeline_parallel.cc | 4 +- .../src/nn/parallel/pp/pipeline_schedule.cc | 2 + infini_train/src/nn/parallel/process_group.cc | 67 +++++ .../src/nn/parallel/tensor_parallel.cc | 33 ++- infini_train/src/nn/parallel/utils.cc | 91 +++++++ infini_train/src/optimizer.cc | 15 +- infini_train/src/tensor.cc | 19 +- .../test_checkpoint_serialization.cc | 86 ++++--- tests/checkpoint/test_lr_scheduler_state.cc | 5 +- tests/checkpoint/test_trainer_state.cc | 26 +- tests/lora/test_lora.cc | 4 +- tests/tensor/test_tensor_copy.cc | 30 +++ .../test_transformer_architecture.cc | 2 +- 60 files changed, 774 insertions(+), 793 deletions(-) create mode 100644 infini_train/include/checkpoint/constants.h delete mode 100644 infini_train/include/checkpoint/reshard.h rename infini_train/include/{checkpoint => }/shard_spec.h (73%) delete mode 100644 infini_train/src/checkpoint/reshard.cc diff --git a/example/gpt2/checkpoint_loader.cc b/example/gpt2/checkpoint_loader.cc index 95e54730b..ceca604be 100644 --- a/example/gpt2/checkpoint_loader.cc +++ b/example/gpt2/checkpoint_loader.cc @@ -163,15 +163,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight - int local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { - auto &tensor - = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), - nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)]; + auto &tensor = state_dict[std::format( + "{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, + std::to_string(idx), nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_w_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_w_bytes, std::ios::cur); @@ -179,14 +176,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.bias - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_b_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_b_bytes, std::ios::cur); @@ -194,13 +189,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_attn.weight (ColumnParallelLinear, but actually applies on "rows") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; // NOTE(zbl): In the .bin model file, Q/K/V is concated along last dim, // i.e. [Q|K|V].T = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn].T // However, each tp_rank needs to get [q_i|k_i|v_i].T, so we need to jump and read them @@ -229,7 +223,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/rows_all, /*cols=*/cols_all, /*row_start=*/2 * n_embd + tp_rank * local_C, /*row_cnt=*/local_C); - ++local_layer_index; } else { size_t c_attn_w_bytes = qkv_out * n_embd * sizeof(float); ifs.seekg(c_attn_w_bytes, std::ios::cur); @@ -237,13 +230,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_attn.bias (ColumnParallelLinear) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamBiasName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamBiasName)]; // NOTE(zbl): Same as c_attn.weight, the bias for Q/K/V is concated // i.e. [Q|K|V] = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn] // However, each tp_rank needs to get [q_i|k_i|v_i], so we need to jump and read them @@ -271,7 +263,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*len=*/len_all, /*start=*/2 * n_embd + tp_rank * local_C, /*cnt=*/local_C); - ++local_layer_index; } else { size_t c_attn_b_bytes = qkv_out * sizeof(float); ifs.seekg(c_attn_b_bytes, std::ios::cur); @@ -279,16 +270,14 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight (RowParallelLinear, but actually applies on "columns") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), n_embd, n_embd, tp_rank * in_pp, in_pp); - ++local_layer_index; } else { size_t c_proj_w_bytes = n_embd * n_embd * sizeof(float); ifs.seekg(c_proj_w_bytes, std::ios::cur); @@ -296,15 +285,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.bias (RowParallelLinear, no shard on bias) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamBiasName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t c_proj_b_bytes = n_embd * sizeof(float); ifs.seekg(c_proj_b_bytes, std::ios::cur); @@ -312,15 +299,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { - auto &tensor - = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), - nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)]; + auto &tensor = state_dict[std::format( + "{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, + std::to_string(idx), nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_w_bytes = n_embd * sizeof(float); ifs.seekg(ln_2_w_bytes, std::ios::cur); @@ -328,14 +312,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.bias - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_b_bytes = n_embd * sizeof(float); ifs.seekg(ln_2_b_bytes, std::ios::cur); @@ -343,15 +325,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (ColumnParallelLinear, but actually applies on "rows") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), fc_out, n_embd, fc_start, fc_pp); - ++local_layer_index; } else { size_t c_fc_w_bytes = fc_out * n_embd * sizeof(float); ifs.seekg(c_fc_w_bytes, std::ios::cur); @@ -359,15 +339,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.bias (ColumnParallelLinear) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamBiasName)]; ReadVectorShardFloat(ifs, static_cast(tensor->DataPtr()), fc_out, fc_start, fc_pp); - ++local_layer_index; } else { size_t c_fc_b_bytes = fc_out * sizeof(float); ifs.seekg(c_fc_b_bytes, std::ios::cur); @@ -375,16 +353,14 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight (RowParallelLinear, but actually applies on "columns") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), n_embd, fc_out, tp_rank * in4_pp, in4_pp); - ++local_layer_index; } else { size_t c_proj_w_bytes = fc_out * n_embd * sizeof(float); ifs.seekg(c_proj_w_bytes, std::ios::cur); @@ -392,15 +368,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.bias (RowParallelLinear, no shard on bias) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t c_proj_b_bytes = n_embd * sizeof(float); ifs.seekg(c_proj_b_bytes, std::ios::cur); diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index ec41c7465..f63c808d6 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -357,7 +357,6 @@ void Train(const nn::parallel::Rank &rank) { } else { optimizer = optimizer_creator(named_parameters); } - const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration; TrainingLRSchedulerConfig sched_config; sched_config.lr = static_cast(FLAGS_learning_rate); @@ -422,7 +421,8 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, @@ -517,6 +517,7 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish backward"; } + nn::parallel::FinalizeModelGrads({model}); optimizer->Step(); if (scheduler) { scheduler->Step(); diff --git a/example/llama3/checkpoint_loader.cc b/example/llama3/checkpoint_loader.cc index 32b11e430..5622afc44 100644 --- a/example/llama3/checkpoint_loader.cc +++ b/example/llama3/checkpoint_loader.cc @@ -181,14 +181,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight : Full version nn::RMSNorm - int local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn1LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_bytes, std::ios::cur); @@ -197,13 +195,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) // transformer.h.{i}.attn.c_attn.weight : ColumnParallelLinear, but actually applies on "rows" // W-qkv should be [Q(=n_embd) | K(=n_kv_head*head_dim) | V(=n_kv_head*head_dim)] × n_embd - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()); const std::streampos base_pos = ifs.tellg(); @@ -229,7 +226,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/attn_rows_all, /*cols=*/attn_cols, /*row_start=*/q_out_rows + kv_out_rows + tp_rank * kv_local_rows, /*row_cnt=*/kv_local_rows); - ++local_layer_index; } else { size_t qkv_bytes = static_cast(attn_rows_all) * attn_cols * sizeof(float); ifs.seekg(qkv_bytes, std::ios::cur); @@ -237,17 +233,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/n_embd, /*col_start=*/tp_rank * in_pp, /*col_cnt=*/in_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * n_embd * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); @@ -255,14 +249,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight : Full version RMSNorm - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn2LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_bytes = static_cast(n_embd) * sizeof(float); ifs.seekg(ln_2_bytes, std::ios::cur); @@ -270,18 +262,16 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (up) -> local packed c_fc rows [fc_pp : 2*fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()) + fc_pp * n_embd; ReadMatrixRowShardFloat(ifs, dst, /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc_bytes, std::ios::cur); @@ -289,17 +279,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc2.weight (gate) -> local packed c_fc rows [0 : fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc2_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc2_bytes, std::ios::cur); @@ -307,17 +295,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/fc_out, /*col_start=*/tp_rank * in_fc_pp, /*col_cnt=*/in_fc_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * ffn_hidden * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 12fecce2b..46d8c8a64 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -347,7 +347,6 @@ void Train(const nn::parallel::Rank &rank) { } else { optimizer = optimizer_creator(named_parameters); } - const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration; TrainingLRSchedulerConfig sched_config; sched_config.lr = static_cast(FLAGS_learning_rate); @@ -412,7 +411,8 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, @@ -504,6 +504,7 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish backward"; } + nn::parallel::FinalizeModelGrads({model}); optimizer->Step(); if (scheduler) { scheduler->Step(); diff --git a/example/mixtral/main.cc b/example/mixtral/main.cc index ce4c8cf36..4998675f8 100644 --- a/example/mixtral/main.cc +++ b/example/mixtral/main.cc @@ -22,6 +22,7 @@ #include "infini_train/include/nn/modules/loss.h" #include "infini_train/include/nn/modules/transformer/transformer.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/utils.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" #ifdef PROFILE_MODE @@ -150,6 +151,7 @@ int main(int argc, char *argv[]) { auto loss_cpu = loss->To(Device()); lossf += static_cast(loss_cpu.DataPtr())[0]; } + infini_train::nn::parallel::FinalizeModelGrads({model}); optimizer->Step(); device_impl->SynchronizeDevice(train_device); diff --git a/example/qwen3/checkpoint_loader.cc b/example/qwen3/checkpoint_loader.cc index 181217a32..7b9b9e65e 100644 --- a/example/qwen3/checkpoint_loader.cc +++ b/example/qwen3/checkpoint_loader.cc @@ -201,35 +201,31 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight : Full version nn::RMSNorm - int local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn1LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_bytes, std::ios::cur); } } - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &q_norm_tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kQNormLayerName, nn::RMSNorm::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kQNormLayerName, + nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(q_norm_tensor->DataPtr()), head_dim); auto &k_norm_tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kKNormLayerName, nn::RMSNorm::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kKNormLayerName, + nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(k_norm_tensor->DataPtr()), head_dim); - ++local_layer_index; } else { size_t qk_norm_bytes = 2 * head_dim * sizeof(float); ifs.seekg(qk_norm_bytes, std::ios::cur); @@ -238,13 +234,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) // transformer.h.{i}.attn.c_attn.weight : ColumnParallelLinear, but actually applies on "rows" // W-qkv should be [Q(=n_embd) | K(=n_kv_head*head_dim) | V(=n_kv_head*head_dim)] x n_embd - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()); const std::streampos base_pos = ifs.tellg(); @@ -270,7 +265,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/attn_rows_all, /*cols=*/attn_cols, /*row_start=*/q_out_rows + kv_out_rows + tp_rank * kv_local_rows, /*row_cnt=*/kv_local_rows); - ++local_layer_index; } else { size_t qkv_bytes = static_cast(attn_rows_all) * attn_cols * sizeof(float); ifs.seekg(qkv_bytes, std::ios::cur); @@ -278,17 +272,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/n_embd, /*col_start=*/tp_rank * in_pp, /*col_cnt=*/in_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * n_embd * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); @@ -296,14 +288,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight : Full version RMSNorm - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn2LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_bytes = static_cast(n_embd) * sizeof(float); ifs.seekg(ln_2_bytes, std::ios::cur); @@ -311,17 +301,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc2.weight (gate) -> local packed c_fc rows [0 : fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc_bytes, std::ios::cur); @@ -329,17 +317,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (up) -> local packed c_fc rows [fc_pp : 2*fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()) + fc_pp * n_embd; ReadMatrixRowShardFloat(ifs, dst, /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc2_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc2_bytes, std::ios::cur); @@ -347,17 +333,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/fc_out, /*col_start=*/tp_rank * in_fc_pp, /*col_cnt=*/in_fc_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * ffn_hidden * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); diff --git a/example/qwen3/main.cc b/example/qwen3/main.cc index ba969fd21..763afa1e7 100644 --- a/example/qwen3/main.cc +++ b/example/qwen3/main.cc @@ -411,11 +411,13 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, @@ -502,6 +504,7 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish backward"; } + nn::parallel::FinalizeModelGrads({model}); optimizer->Step(); if (scheduler) { scheduler->Step(); diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index d93ea35f3..9703cdc7b 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -9,8 +9,8 @@ #include #include "infini_train/include/checkpoint/save_planner.h" -#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/lr_scheduler.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Optimizer; @@ -27,7 +27,8 @@ struct TrainerState { int64_t n_head = 0; int64_t n_kv_head = 0; int64_t n_embd = 0; - int64_t vocab_size = 0; + int64_t original_vocab_size = 0; + int64_t padded_vocab_size = 0; int ddp_size = 1; int tp_size = 1; int sp_size = 1; @@ -43,12 +44,6 @@ class Checkpoint { static void Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, LRScheduler *lr_scheduler); - static void SaveSharded(const std::filesystem::path &checkpoint_dir, const checkpoint::ShardedStateDict &sharded_sd, - const std::vector &write_items, - const std::unordered_map> &state_dict, - const std::unordered_map> &optimizer_state, - const TrainerState &state, int global_rank); - static void SaveStateDictFile(const std::filesystem::path &path, const std::unordered_map> &state_dict); @@ -57,15 +52,6 @@ class Checkpoint { struct CheckpointMetadata { int version = 0; - int64_t iteration = 0; - - struct ParallelConfig { - int tp_size = 1; - int pp_size = 1; - int dp_size = 1; - int sp_size = 1; - int vpp_size = 1; - } parallel_config; struct TensorEntry { std::string key; @@ -74,7 +60,7 @@ class Checkpoint { std::vector local_shape; std::vector global_offset; std::vector axis_fragmentations; - std::vector segments; + std::vector segments; std::string file; uint64_t offset = 0; uint64_t byte_size = 0; @@ -89,14 +75,6 @@ class Checkpoint { static CheckpointMetadata LoadMetadata(const std::filesystem::path &checkpoint_dir); static void SaveMetadataFile(const std::filesystem::path &path, const CheckpointMetadata &metadata); - // Public LR-scheduler serialization helpers used by checkpoint_manager. - static void SaveLRSchedulerStateFile(const std::filesystem::path &path, const LRSchedulerStateDict &state_dict); - static LRSchedulerStateDict LoadLRSchedulerStateFile(const std::filesystem::path &path); - - // Public trainer-state serialization helpers used by checkpoint_manager. - static void SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state); - static TrainerState LoadTrainerStateFile(const std::filesystem::path &path); - private: struct SavedTensorLocation { uint64_t data_offset = 0; @@ -110,6 +88,12 @@ class Checkpoint { static std::unordered_map> LoadStateDict(const std::filesystem::path &path); + static void SaveLocalShard(const std::filesystem::path &checkpoint_dir, const ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + int global_rank); + static void SaveTrainerState(const std::filesystem::path &path, const TrainerState &state); static TrainerState LoadTrainerState(const std::filesystem::path &path); }; diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index d90bf655f..fa4a3cc61 100644 --- a/infini_train/include/checkpoint/checkpoint_manager.h +++ b/infini_train/include/checkpoint/checkpoint_manager.h @@ -44,7 +44,8 @@ struct SaveCheckpointArgs { int64_t n_head = 0; int64_t n_kv_head = 0; int64_t n_embd = 0; - int64_t vocab_size = 0; + int64_t original_vocab_size = 0; + int64_t padded_vocab_size = 0; int ddp_size = 1; int tp_size = 1; int sp_size = 1; diff --git a/infini_train/include/checkpoint/constants.h b/infini_train/include/checkpoint/constants.h new file mode 100644 index 000000000..e7ad0869c --- /dev/null +++ b/infini_train/include/checkpoint/constants.h @@ -0,0 +1,14 @@ +#pragma once + +namespace infini_train::checkpoint { + +inline constexpr char kModelCheckpointFilename[] = "model.ckpt"; +inline constexpr char kOptimizerCheckpointFilename[] = "optimizer.ckpt"; +inline constexpr char kMetadataFilename[] = "metadata.json"; +inline constexpr char kTemporaryMetadataFilename[] = "metadata.json.tmp"; +inline constexpr char kTrainerStateFilename[] = "trainer_state.json"; +inline constexpr char kLRSchedulerFilename[] = "lr_scheduler.ckpt"; +inline constexpr char kLatestIterationFilename[] = "latest_checkpointed_iteration.txt"; +inline constexpr char kTemporaryLatestIterationFilename[] = "latest_checkpointed_iteration.txt.tmp"; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/load_planner.h b/infini_train/include/checkpoint/load_planner.h index 7011cc730..bbca3fc75 100644 --- a/infini_train/include/checkpoint/load_planner.h +++ b/infini_train/include/checkpoint/load_planner.h @@ -6,8 +6,8 @@ #include #include "infini_train/include/checkpoint/checkpoint.h" -#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/datatype.h" +#include "infini_train/include/shard_spec.h" namespace infini_train::checkpoint { diff --git a/infini_train/include/checkpoint/reshard.h b/infini_train/include/checkpoint/reshard.h deleted file mode 100644 index c091d2757..000000000 --- a/infini_train/include/checkpoint/reshard.h +++ /dev/null @@ -1,22 +0,0 @@ -#pragma once - -#include - -#include "infini_train/include/checkpoint/checkpoint.h" - -namespace infini_train { -class LRScheduler; -class Optimizer; -namespace nn { -class Module; -} -} // namespace infini_train - -namespace infini_train::checkpoint { - -// Restore this rank's target model shards from a distributed checkpoint. -void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, - TrainerState &state, LRScheduler *lr_scheduler, - const Checkpoint::CheckpointMetadata &metadata); - -} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/save_planner.h b/infini_train/include/checkpoint/save_planner.h index 02a5a21c1..3c5a6fe3c 100644 --- a/infini_train/include/checkpoint/save_planner.h +++ b/infini_train/include/checkpoint/save_planner.h @@ -6,8 +6,8 @@ #include #include -#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/datatype.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Tensor; @@ -18,8 +18,7 @@ namespace infini_train::checkpoint { // Physical write description for one local tensor shard. struct WriteItem { std::string key; - std::string filename; // "model.ckpt" or "optimizer.ckpt" - uint64_t offset = 0; // Planned byte offset in the checkpoint file. + std::string filename; // Checkpoint payload filename. uint64_t byte_size = 0; // Tensor payload size in bytes. DataType dtype = DataType::kFLOAT32; std::vector local_shape; @@ -42,38 +41,7 @@ BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, inline uint64_t TensorByteSize(DataType dtype, const std::vector &shape) { uint64_t numel = 1; for (auto d : shape) { numel *= static_cast(d); } - switch (dtype) { - case DataType::kBFLOAT16: - case DataType::kFLOAT16: - return numel * 2; - case DataType::kFLOAT32: - return numel * 4; - case DataType::kFLOAT64: - case DataType::kINT64: - case DataType::kUINT64: - return numel * 8; - case DataType::kINT32: - case DataType::kUINT32: - return numel * 4; - case DataType::kINT16: - case DataType::kUINT16: - return numel * 2; - case DataType::kINT8: - case DataType::kUINT8: - case DataType::kBOOL: - return numel; - default: - return numel * 4; - } -} - -// Compute one rank's balanced interval, including non-divisible dimensions. -inline std::pair GetRankSliceRange(int64_t global_size, int world_size, int rank) { - int64_t per_rank = global_size / world_size; - int64_t remainder = global_size % world_size; - int64_t start = rank * per_rank + std::min(rank, remainder); - int64_t local_size = per_rank + (rank < remainder ? 1 : 0); - return {start, local_size}; + return numel * static_cast(kDataTypeToSize.at(dtype)); } } // namespace infini_train::checkpoint diff --git a/infini_train/include/core/ccl/ccl.h b/infini_train/include/core/ccl/ccl.h index 899a12fcd..47b9123a2 100644 --- a/infini_train/include/core/ccl/ccl.h +++ b/infini_train/include/core/ccl/ccl.h @@ -54,6 +54,9 @@ class CclImpl { nn::parallel::function::ReduceOpType reduce_op, const CclComm *comm, Stream *stream) const; + virtual void AlltoAll(const void *sendbuff, void *recvbuff, size_t count, DataType dtype, const CclComm *comm, + Stream *stream) const; + virtual void Send(const void *buff, size_t count, DataType dtype, int peer, const CclComm *comm, Stream *stream) const; diff --git a/infini_train/include/nn/lora/lora_parallel_linear.h b/infini_train/include/nn/lora/lora_parallel_linear.h index b485e3c7c..5c9bbddba 100644 --- a/infini_train/include/nn/lora/lora_parallel_linear.h +++ b/infini_train/include/nn/lora/lora_parallel_linear.h @@ -34,7 +34,7 @@ class LoRAColumnParallelLinear : public nn::parallel::ColumnParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; void MergeWeights(); void UnmergeWeights(); @@ -76,7 +76,7 @@ class LoRARowParallelLinear : public nn::parallel::RowParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; void MergeWeights(); void UnmergeWeights(); diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 0816d247f..40c5f83ac 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -6,9 +6,9 @@ #include #include -#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/datatype.h" #include "infini_train/include/device.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Tensor; @@ -79,7 +79,7 @@ class Module : public std::enable_shared_from_this { virtual std::unordered_map> StateDict() const; // Return state-dict metadata with global shard coordinates. - virtual checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const; + virtual ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const; // Current behavior: missing keys / shape / dtype mismatches are FATAL errors; unexpected keys in state_dict are // WARNING-only and silently ignored. diff --git a/infini_train/include/nn/modules/transformer/causal_self_attention.h b/infini_train/include/nn/modules/transformer/causal_self_attention.h index 6e713b19f..15f2f8b1e 100644 --- a/infini_train/include/nn/modules/transformer/causal_self_attention.h +++ b/infini_train/include/nn/modules/transformer/causal_self_attention.h @@ -25,7 +25,7 @@ class CausalSelfAttention : public infini_train::nn::CloneableModule> Forward(const std::vector> &x) override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; private: TransformerConfig config_; diff --git a/infini_train/include/nn/modules/transformer/transformer.h b/infini_train/include/nn/modules/transformer/transformer.h index 455f37c28..40990cb9f 100644 --- a/infini_train/include/nn/modules/transformer/transformer.h +++ b/infini_train/include/nn/modules/transformer/transformer.h @@ -78,10 +78,9 @@ class TransformerModel : public CloneableModule { const TransformerConfig &Config() const { return config_; } - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; - void LoadStateDict(const std::unordered_map> &state_dict) override; private: const TransformerConfig config_; diff --git a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h index 4aa130dc8..45edfd91d 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -34,11 +34,13 @@ class DistributedDataParallel : public nn::Module { std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; std::unordered_map> StateDict() const override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; void LoadStateDict(const std::unordered_map> &state_dict) override; std::unique_ptr no_sync() override; + void FinishGradSync(); + DistributedDataParallelConfig ddp_config() const { return ddp_config_; } const std::vector> ¶m_grad_buffers() const { return param_grad_buffers_; } diff --git a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h index d7cea198e..b27e6f7b0 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_optimizer.h +++ b/infini_train/include/nn/parallel/ddp/distributed_optimizer.h @@ -37,9 +37,6 @@ class DistributedOptimizer final : public infini_train::Optimizer { void LoadStateDict(const std::unordered_map> &state_dict) override; - void StartGradSync(); - void FinishGradSync(); - void StartParamSync(bool force_sync = false); void FinishParamSync(bool skip_next_bucket_dispatch = false); diff --git a/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h b/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h index 2c572984d..a7866c4b3 100644 --- a/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h +++ b/infini_train/include/nn/parallel/ddp/param_and_grad_buffer.h @@ -5,6 +5,7 @@ #include #include #include +#include #include #include "infini_train/include/datatype.h" @@ -19,6 +20,9 @@ class Work; } // namespace infini_train namespace infini_train::nn::parallel { +// Original parameter and its local gradient view, registered by DistributedOptimizer. +using LocalGradShard = std::pair, std::shared_ptr>; + class ParamAndGradBucket { public: /** @@ -123,6 +127,10 @@ class ParamAndGradBucketGroup { // ZeRO-2: Get a bucket's local grad shard buffer std::shared_ptr GetLocalGradShardBuffer(size_t bucket_idx) const; + const std::vector &local_grad_shards() const; + + void set_local_grad_shards(std::vector shards); + const DistributedDataParallelConfig &config() const { return ddp_config_; } private: @@ -146,6 +154,8 @@ class ParamAndGradBucketGroup { // ZeRO-2: persistent grad shard buffers and temporary full grad buffers std::vector> grad_shard_buffer_list_; std::vector> temp_full_grad_buffer_list_; + // These views share optimizer grad storage and persist across iteration resets. + std::vector local_grad_shards_; std::shared_ptr next_param_gather_bucket_group_ = nullptr; diff --git a/infini_train/include/nn/parallel/parallel_functional.h b/infini_train/include/nn/parallel/parallel_functional.h index 2eed56f48..a5641d34b 100644 --- a/infini_train/include/nn/parallel/parallel_functional.h +++ b/infini_train/include/nn/parallel/parallel_functional.h @@ -25,6 +25,9 @@ std::shared_ptr AllGather(const std::shared_ptr &output, const std std::shared_ptr ReduceScatter(const std::shared_ptr &output, const std::shared_ptr &input, ReduceOpType reduce_op, const ProcessGroup *pg = nullptr, bool async_op = false); +std::shared_ptr AlltoAll(const std::shared_ptr &output, const std::shared_ptr &input, + const ProcessGroup *pg = nullptr, bool async_op = false); + std::vector>> Scatter(const std::vector> &input_tensors, const std::vector &device_ids, int dim); diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 58f48cd59..ee717d4ef 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_parallel.h +++ b/infini_train/include/nn/parallel/pp/pipeline_parallel.h @@ -43,7 +43,7 @@ class PipelineParallel : public Module { std::unordered_map> StateDict() const override; std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; void LoadStateDict(const std::unordered_map> &state_dict) override; private: diff --git a/infini_train/include/nn/parallel/process_group.h b/infini_train/include/nn/parallel/process_group.h index 2dccc33e0..ae93ec246 100644 --- a/infini_train/include/nn/parallel/process_group.h +++ b/infini_train/include/nn/parallel/process_group.h @@ -30,6 +30,17 @@ class Work; namespace infini_train::nn::parallel { +enum class P2POpType { + kSend, + kRecv, +}; + +struct P2POp { + P2POpType type; + std::shared_ptr tensor; + int peer_rank; +}; + class ProcessGroup { public: explicit ProcessGroup(Device::DeviceType backend, const std::string &process_group_name, @@ -63,12 +74,17 @@ class ProcessGroup { const std::vector> &input_tensors, int root_rank_in_group, bool async_op = false) const; + virtual std::shared_ptr AlltoAll(const std::shared_ptr &output, const std::shared_ptr &input, + bool async_op = false) const; + virtual std::shared_ptr Send(std::vector> tensors, int dest_rank, bool async_op = false) const; virtual std::shared_ptr Recv(std::vector> tensors, int src_rank, bool async_op = false) const; + virtual std::shared_ptr BatchSendRecv(const std::vector &ops, bool async_op = false) const; + // Legacy communication APIs (Single-stream) // FIXME(dcj): BroadCast_ and Scatter_ are temporarily retained with trailing underscores for existing DP callers. // Replace direct DP usage with a higher-level communication abstraction. diff --git a/infini_train/include/nn/parallel/tensor_parallel.h b/infini_train/include/nn/parallel/tensor_parallel.h index 4fbedb2a5..d5049877a 100644 --- a/infini_train/include/nn/parallel/tensor_parallel.h +++ b/infini_train/include/nn/parallel/tensor_parallel.h @@ -4,9 +4,9 @@ #include #include "infini_train/include/autograd/function.h" -#include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Tensor; @@ -38,7 +38,7 @@ class ColumnParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; protected: bool bias_ = true; @@ -69,7 +69,7 @@ class RowParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; protected: bool bias_ = true; @@ -90,7 +90,7 @@ class VocabParallelEmbedding : public nn::CloneableModule> Forward(const std::vector> &input_tensors) override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; private: bool reduce_scatter_embeddings_ = false; // whether to perform ReduceScatter after embedding lookup diff --git a/infini_train/include/nn/parallel/utils.h b/infini_train/include/nn/parallel/utils.h index 8aa11856a..1d3fe975c 100644 --- a/infini_train/include/nn/parallel/utils.h +++ b/infini_train/include/nn/parallel/utils.h @@ -7,6 +7,9 @@ namespace infini_train { class Tensor; +namespace nn { +class Module; +} } // namespace infini_train namespace infini_train::nn::parallel { @@ -31,4 +34,6 @@ std::vector> GatherFromSPRegionFunc(const std::shared_pt std::vector> ScatterToTPRegionFunc(const std::shared_ptr &input); std::vector> ReduceFromTPRegionFunc(const std::shared_ptr &input); std::vector> CopyToTPRegionFunc(const std::shared_ptr &input); + +void FinalizeModelGrads(const std::vector> &model_chunks); } // namespace infini_train::nn::parallel diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index d85b1acea..377924fc3 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -4,6 +4,7 @@ #include #include #include +#include #include #include #include @@ -52,6 +53,11 @@ class Optimizer { }; namespace optimizers { +inline constexpr std::string_view kAdamOptimizerPrefix = "adam."; +inline constexpr std::string_view kAdamFirstMomentPrefix = "adam.m."; +inline constexpr std::string_view kAdamSecondMomentPrefix = "adam.v."; +inline constexpr std::string_view kAdamStepKey = "adam.t"; + class SGD : public Optimizer { public: SGD(const std::vector> ¶ms, float learning_rate); diff --git a/infini_train/include/checkpoint/shard_spec.h b/infini_train/include/shard_spec.h similarity index 73% rename from infini_train/include/checkpoint/shard_spec.h rename to infini_train/include/shard_spec.h index 75263765a..b5af5d24d 100644 --- a/infini_train/include/checkpoint/shard_spec.h +++ b/infini_train/include/shard_spec.h @@ -3,13 +3,14 @@ #include #include #include +#include #include #include "glog/logging.h" #include "infini_train/include/datatype.h" -namespace infini_train::checkpoint { +namespace infini_train { struct ShardSegment { int64_t global_offset = 0; @@ -28,6 +29,7 @@ struct ShardedTensor { std::vector local_shape; std::vector global_offset; std::vector axis_fragmentations; + bool allow_shape_mismatch = false; // Optional disjoint regions along the single fragmented axis. This is used // by layouts such as rank-local [Q, K, V], which are not one contiguous // slice of the logical global [Q, K, V] tensor. @@ -37,10 +39,20 @@ struct ShardedTensor { return key == other.key && local_key == other.local_key && dtype == other.dtype && global_shape == other.global_shape && local_shape == other.local_shape && global_offset == other.global_offset && axis_fragmentations == other.axis_fragmentations - && segments == other.segments; + && allow_shape_mismatch == other.allow_shape_mismatch && segments == other.segments; } }; +inline ShardedTensor MakeShardedTensor(std::string key, DataType dtype, std::vector shape) { + const auto ndim = shape.size(); + return {.key = std::move(key), + .dtype = dtype, + .global_shape = shape, + .local_shape = std::move(shape), + .global_offset = std::vector(ndim, 0), + .axis_fragmentations = std::vector(ndim, 1)}; +} + struct ShardedStateDict { std::map tensors; @@ -53,4 +65,4 @@ struct ShardedStateDict { } }; -} // namespace infini_train::checkpoint +} // namespace infini_train diff --git a/infini_train/include/tensor.h b/infini_train/include/tensor.h index dcfd8927f..3aabe28f1 100644 --- a/infini_train/include/tensor.h +++ b/infini_train/include/tensor.h @@ -80,6 +80,9 @@ class Tensor : public std::enable_shared_from_this { size_t NumElements() const; DataType Dtype() const; + void set_sequence_parallel(bool enabled); + bool sequence_parallel() const; + std::shared_ptr Detach() const; void Fill(Scalar value); @@ -242,6 +245,7 @@ class Tensor : public std::enable_shared_from_this { private: std::shared_ptr grad_ = nullptr; bool requires_grad_ = false; + bool sequence_parallel_ = false; bool is_leaf_ = true; std::shared_ptr grad_fn_ = nullptr; int output_idx_ = 0; diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index b76f5b410..a291ec698 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -3,6 +3,7 @@ #include #include #include +#include #include #include #include @@ -12,10 +13,15 @@ #include "glog/logging.h" +#include "infini_train/include/checkpoint/constants.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/load_strategy.h" #include "infini_train/include/checkpoint/save_planner.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/parallel_functional.h" +#include "infini_train/include/nn/parallel/work.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -182,67 +188,100 @@ template T ExtractNumberField(const std::string &content, const std } return value; } + +void SynchronizeCheckpointRanks(const nn::Module &model) { + if (nn::parallel::global::GetWorldSize() == 1) { + return; + } + const auto parameters = model.Parameters(); + CHECK(!parameters.empty()) << "Cannot synchronize checkpoint save for a model without parameters"; + auto token = std::make_shared(std::vector{1}, DataType::kFLOAT32, parameters.front()->GetDevice()); + token->Fill(1.0f); + nn::parallel::function::AllReduce(token, nn::parallel::function::ReduceOpType::kSum, nullptr, true)->Synchronize(); +} } // namespace void Checkpoint::Save(const std::filesystem::path &checkpoint_dir, const nn::Module &model, const Optimizer *optimizer, const TrainerState &state, const LRScheduler *lr_scheduler) { std::filesystem::create_directories(checkpoint_dir); - LOG(INFO) << "[CKPT] Save begin: dir=" << checkpoint_dir << ", global_step=" << state.global_step; - - const auto model_path = checkpoint_dir / ("model.ckpt"); - - SaveStateDict(model_path, model.StateDict()); - - if (optimizer != nullptr) { - auto opt_state = optimizer->StateDict(); - if (!opt_state.empty()) { - const auto opt_path = checkpoint_dir / "optimizer.ckpt"; - SaveStateDict(opt_path, opt_state); + const int global_rank = nn::parallel::global::thread_global_rank; + const auto staging_root = checkpoint_dir / ".metadata_tmp"; + if (global_rank == 0) { + std::filesystem::remove_all(staging_root); + } + SynchronizeCheckpointRanks(model); + + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(global_rank, dp_rank, tp_rank, pp_rank); + // TODO(jym): Select checkpoint writers from each shard's replica_id instead of hard-coding DP rank 0. + if (dp_rank == 0) { + const auto rank_dir = checkpoint_dir / std::format("rank_{:06d}", global_rank); + auto sharded_state = model.BuildShardedStateDict(); + std::unordered_map> optimizer_state; + if (optimizer != nullptr) { + optimizer_state = optimizer->StateDict(); + sharded_state.Merge(checkpoint::BuildOptimizerShardedStateDict(sharded_state, optimizer_state)); } - } + const auto write_items = checkpoint::SavePlanner::Plan(sharded_state, global_rank); + SaveLocalShard(rank_dir, sharded_state, write_items, model.StateDict(), optimizer_state, global_rank); - if (lr_scheduler != nullptr) { - SaveLRSchedulerState(checkpoint_dir / "lr_scheduler.ckpt", lr_scheduler->StateDict()); + const auto staging_rank_dir = staging_root / std::format("rank_{:06d}", global_rank); + std::filesystem::create_directories(staging_rank_dir); + std::filesystem::rename(rank_dir / checkpoint::kMetadataFilename, + staging_rank_dir / checkpoint::kMetadataFilename); } - SaveTrainerState(checkpoint_dir / "trainer_state.json", state); - LOG(ERROR) << "[CKPT] Save done: dir=" << checkpoint_dir; + SynchronizeCheckpointRanks(model); + if (global_rank == 0) { + SaveTrainerState(checkpoint_dir / checkpoint::kTrainerStateFilename, state); + if (lr_scheduler != nullptr) { + SaveLRSchedulerState(checkpoint_dir / checkpoint::kLRSchedulerFilename, lr_scheduler->StateDict()); + } + auto metadata = LoadMetadata(staging_root); + CHECK(metadata.has_metadata); + SaveMetadataFile(checkpoint_dir / checkpoint::kTemporaryMetadataFilename, metadata); + if (std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) { + std::filesystem::remove(checkpoint_dir / checkpoint::kMetadataFilename); + } + std::filesystem::rename(checkpoint_dir / checkpoint::kTemporaryMetadataFilename, + checkpoint_dir / checkpoint::kMetadataFilename); + std::filesystem::remove_all(staging_root); + } + SynchronizeCheckpointRanks(model); } void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, LRScheduler *lr_scheduler) { - const auto model_path = checkpoint_dir / "model.ckpt"; - LOG(INFO) << "[CKPT] Loading model: " << model_path; - - model.LoadStateDict(LoadStateDict(model_path)); + const auto metadata = LoadMetadata(checkpoint_dir); + CHECK(metadata.has_metadata); + CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + state = LoadTrainerState(checkpoint_dir / checkpoint::kTrainerStateFilename); + // TODO(jym): Support VPP checkpoint resharding by describing virtual pipeline chunks in the target shard layout. + CHECK_EQ(state.vpp_size, 1) << "Checkpoint resharding with saved VPP is not supported yet"; + CHECK_EQ(nn::parallel::global::GetVirtualPipelineParallelSize(), 1) + << "Checkpoint resharding with runtime VPP is not supported yet"; + + auto model_sharded_state = model.BuildShardedStateDict(); + checkpoint::IndexedRegionLoadStrategy strategy; + auto result = strategy.Execute(checkpoint_dir, checkpoint::LoadPlanner::PlanReshard(metadata, model_sharded_state)); + model.LoadStateDict(result); + + const int current_tp = nn::parallel::global::GetTensorParallelSize(); + const int current_pp = nn::parallel::global::GetPipelineParallelSize(); + state.tp_size = current_tp; + state.pp_size = current_pp; + state.ddp_size = nn::parallel::global::GetDataParallelSize(); + state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; if (optimizer != nullptr) { - const auto opt_path = checkpoint_dir / "optimizer.ckpt"; - if (std::filesystem::exists(opt_path)) { - LOG(INFO) << "[CKPT] Loading optimizer: " << opt_path; - optimizer->LoadStateDict(LoadStateDict(opt_path)); - } else { - LOG(FATAL) << "Optimizer checkpoint not found at: " << opt_path; - } + auto optimizer_sharded_state + = checkpoint::BuildOptimizerShardedStateDict(model_sharded_state, optimizer->StateDict()); + auto optimizer_plan = checkpoint::LoadPlanner::PlanReshard(metadata, optimizer_sharded_state); + optimizer->LoadStateDict(strategy.Execute(checkpoint_dir, optimizer_plan)); } - - state = LoadTrainerState(checkpoint_dir / "trainer_state.json"); - - if (lr_scheduler != nullptr) { - const auto lr_scheduler_path = checkpoint_dir / "lr_scheduler.ckpt"; - if (std::filesystem::exists(lr_scheduler_path)) { - LOG(INFO) << "[CKPT] Loading LR scheduler: " << lr_scheduler_path; - lr_scheduler->LoadStateDict(LoadLRSchedulerState(lr_scheduler_path)); - } else { - LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found at: " << lr_scheduler_path - << ". Keeping the initialized scheduler state."; - } + if (lr_scheduler != nullptr && std::filesystem::exists(checkpoint_dir / checkpoint::kLRSchedulerFilename)) { + lr_scheduler->LoadStateDict(LoadLRSchedulerState(checkpoint_dir / checkpoint::kLRSchedulerFilename)); } - - LOG(ERROR) << "[CKPT] Load done: global_step=" << state.global_step - << ", consumed_train_samples=" << state.consumed_train_samples << ", topology(ddp,tp,sp,pp)=(" - << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << "," - << state.vpp_size << ")"; } Checkpoint::SavedTensorLocations @@ -330,7 +369,8 @@ void Checkpoint::SaveTrainerState(const std::filesystem::path &path, const Train ofs << " \"n_head\": " << state.n_head << ",\n"; ofs << " \"n_kv_head\": " << state.n_kv_head << ",\n"; ofs << " \"n_embd\": " << state.n_embd << ",\n"; - ofs << " \"vocab_size\": " << state.vocab_size << ",\n"; + ofs << " \"original_vocab_size\": " << state.original_vocab_size << ",\n"; + ofs << " \"padded_vocab_size\": " << state.padded_vocab_size << ",\n"; ofs << " \"global_step\": " << state.global_step << ",\n"; ofs << " \"consumed_train_samples\": " << state.consumed_train_samples << ",\n"; ofs << " \"ddp_size\": " << state.ddp_size << ",\n"; @@ -352,7 +392,9 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.n_head = ExtractNumberField(content, "n_head", 0); state.n_kv_head = ExtractNumberField(content, "n_kv_head", 0); state.n_embd = ExtractNumberField(content, "n_embd", 0); - state.vocab_size = ExtractNumberField(content, "vocab_size", 0); + const auto legacy_vocab_size = ExtractNumberField(content, "vocab_size", 0); + state.original_vocab_size = ExtractNumberField(content, "original_vocab_size", 0); + state.padded_vocab_size = ExtractNumberField(content, "padded_vocab_size", legacy_vocab_size); state.global_step = ExtractNumberField(content, "global_step", 0); state.consumed_train_samples = ExtractNumberField(content, "consumed_train_samples", 0); state.ddp_size = ExtractNumberField(content, "ddp_size", 1); @@ -363,20 +405,6 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { return state; } -void Checkpoint::SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state) { - SaveTrainerState(path, state); -} - -TrainerState Checkpoint::LoadTrainerStateFile(const std::filesystem::path &path) { return LoadTrainerState(path); } - -void Checkpoint::SaveLRSchedulerStateFile(const std::filesystem::path &path, const LRSchedulerStateDict &state_dict) { - SaveLRSchedulerState(path, state_dict); -} - -LRSchedulerStateDict Checkpoint::LoadLRSchedulerStateFile(const std::filesystem::path &path) { - return LoadLRSchedulerState(path); -} - void Checkpoint::SaveStateDictFile(const std::filesystem::path &path, const std::unordered_map> &state_dict) { SaveStateDict(path, state_dict); @@ -387,10 +415,6 @@ Checkpoint::LoadStateDictFile(const std::filesystem::path &path) { return LoadStateDict(path); } -// ----------------------------------------------------------------------------- -// Save local shards and a temporary rank manifest from a ShardedStateDict. -// ----------------------------------------------------------------------------- - static std::string DataTypeToString(DataType dt) { auto it = kDataTypeToDesc.find(dt); if (it != kDataTypeToDesc.end()) { @@ -399,15 +423,13 @@ static std::string DataTypeToString(DataType dt) { return "fp32"; } -void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, - const checkpoint::ShardedStateDict &sharded_sd, - const std::vector &write_items, - const std::unordered_map> &state_dict, - const std::unordered_map> &optimizer_state, - const TrainerState &state, int global_rank) { +void Checkpoint::SaveLocalShard(const std::filesystem::path &checkpoint_dir, const ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + int global_rank) { std::filesystem::create_directories(checkpoint_dir); - LOG(INFO) << "[CKPT] SaveSharded begin: dir=" << checkpoint_dir << ", global_step=" << state.global_step - << ", rank=" << global_rank; + LOG(INFO) << "[CKPT] SaveLocalShard begin: dir=" << checkpoint_dir << ", rank=" << global_rank; SavedTensorLocations model_file_index; SavedTensorLocations optimizer_file_index; @@ -417,7 +439,7 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, std::unordered_map> filtered_sd; for (const auto &[key, info] : sharded_sd.tensors) { // Optimizer tensors are serialized separately. - if (key.starts_with("adam.")) { + if (key.starts_with(optimizers::kAdamOptimizerPrefix)) { continue; } // Match metadata keys to the local tensor payloads. @@ -428,38 +450,25 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, } } if (!filtered_sd.empty()) { - model_file_index = SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); + model_file_index = SaveStateDict(checkpoint_dir / checkpoint::kModelCheckpointFilename, filtered_sd); } } // Save the rank-local optimizer state. if (!optimizer_state.empty()) { - optimizer_file_index = SaveStateDict(checkpoint_dir / "optimizer.ckpt", optimizer_state); + optimizer_file_index + = SaveStateDict(checkpoint_dir / checkpoint::kOptimizerCheckpointFilename, optimizer_state); } // Write the temporary rank manifest. { - std::ofstream ofs(checkpoint_dir / "metadata.json"); - CHECK(ofs.is_open()) << "Failed to open metadata.json: " << checkpoint_dir / "metadata.json"; + std::ofstream ofs(checkpoint_dir / checkpoint::kMetadataFilename); + CHECK(ofs.is_open()) << "Failed to open metadata.json: " << checkpoint_dir / checkpoint::kMetadataFilename; ofs << "{\n"; ofs << " \"version\": 3,\n"; ofs << " \"format\": \"infinitrain_sharded\",\n"; - ofs << " \"iteration\": " << state.global_step << ",\n"; - ofs << " \"parallel_config\": {\n"; - ofs << " \"tp_size\": " << state.tp_size << ",\n"; - ofs << " \"pp_size\": " << state.pp_size << ",\n"; - ofs << " \"dp_size\": " << state.ddp_size << ",\n"; - ofs << " \"sp_size\": " << state.sp_size << ",\n"; - ofs << " \"vpp_size\": " << state.vpp_size << "\n"; - ofs << " },\n"; - ofs << " \"model_config\": {\n"; - ofs << " \"n_layer\": " << state.n_layer << ",\n"; - ofs << " \"n_head\": " << state.n_head << ",\n"; - ofs << " \"n_kv_head\": " << state.n_kv_head << ",\n"; - ofs << " \"n_embd\": " << state.n_embd << ",\n"; - ofs << " \"vocab_size\": " << state.vocab_size << "\n"; - ofs << " },\n"; + ofs << " \"tensors\": [\n"; std::vector emitted_items; @@ -473,7 +482,8 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, for (size_t i = 0; i < emitted_items.size(); ++i) { const auto &item = *emitted_items[i]; const auto it = sharded_sd.tensors.find(item.key); - const auto &file_index = item.filename == "optimizer.ckpt" ? optimizer_file_index : model_file_index; + const auto &file_index + = item.filename == checkpoint::kOptimizerCheckpointFilename ? optimizer_file_index : model_file_index; const auto storage_it = file_index.find(item.key); CHECK(storage_it != file_index.end()) << "Missing stored tensor metadata for " << item.key; const auto &storage = storage_it->second; @@ -511,9 +521,9 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, } ofs << "],\n"; }; - write_segments("segment_global_offsets", &checkpoint::ShardSegment::global_offset); - write_segments("segment_local_offsets", &checkpoint::ShardSegment::local_offset); - write_segments("segment_lengths", &checkpoint::ShardSegment::length); + write_segments("segment_global_offsets", &ShardSegment::global_offset); + write_segments("segment_local_offsets", &ShardSegment::local_offset); + write_segments("segment_lengths", &ShardSegment::length); ofs << " \"file\": \"" << item.filename << "\",\n"; ofs << " \"offset\": " << storage.data_offset << ",\n"; @@ -533,7 +543,7 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, LOG(INFO) << "[CKPT] metadata.json written"; } - LOG(ERROR) << "[CKPT] SaveSharded done: dir=" << checkpoint_dir; + LOG(ERROR) << "[CKPT] SaveLocalShard done: dir=" << checkpoint_dir; } // Load one manifest or aggregate writer manifests while finalizing a checkpoint. @@ -556,7 +566,7 @@ static std::string ExtractJsonString(const std::string &obj, const std::string & static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem::path &checkpoint_dir) { Checkpoint::CheckpointMetadata meta; - auto metadata_path = checkpoint_dir / "metadata.json"; + auto metadata_path = checkpoint_dir / checkpoint::kMetadataFilename; if (!std::filesystem::exists(metadata_path)) { meta.has_metadata = false; return meta; @@ -568,12 +578,6 @@ static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem:: meta.has_metadata = true; meta.version = ExtractNumberField(content, "version", 0); - meta.iteration = ExtractNumberField(content, "iteration", 0); - meta.parallel_config.tp_size = ExtractNumberField(content, "tp_size", 1); - meta.parallel_config.pp_size = ExtractNumberField(content, "pp_size", 1); - meta.parallel_config.dp_size = ExtractNumberField(content, "dp_size", 1); - meta.parallel_config.sp_size = ExtractNumberField(content, "sp_size", 1); - meta.parallel_config.vpp_size = ExtractNumberField(content, "vpp_size", 1); // Locate the tensors array. auto tensors_key = content.find("\"tensors\""); @@ -733,19 +737,19 @@ static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem:: obj_pos = obj_end + 1; } - LOG(INFO) << "[CKPT] Loaded metadata.json: " << meta.tensors.size() << " tensors, iteration=" << meta.iteration; + LOG(INFO) << "[CKPT] Loaded metadata.json: " << meta.tensors.size() << " tensor shards"; return meta; } Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::path &checkpoint_dir) { - if (std::filesystem::exists(checkpoint_dir / "metadata.json")) { + if (std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) { return LoadSingleMetadata(checkpoint_dir); } CheckpointMetadata merged; for (const auto &entry : std::filesystem::directory_iterator(checkpoint_dir)) { if (!entry.is_directory() || !entry.path().filename().string().starts_with("rank_") - || !std::filesystem::exists(entry.path() / "metadata.json")) { + || !std::filesystem::exists(entry.path() / checkpoint::kMetadataFilename)) { continue; } auto rank_metadata = LoadSingleMetadata(entry.path()); @@ -771,14 +775,7 @@ void Checkpoint::SaveMetadataFile(const std::filesystem::path &path, const Check ofs << "{\n"; ofs << " \"version\": 3,\n"; ofs << " \"format\": \"infinitrain_sharded\",\n"; - ofs << " \"iteration\": " << metadata.iteration << ",\n"; - ofs << " \"parallel_config\": {\n"; - ofs << " \"tp_size\": " << metadata.parallel_config.tp_size << ",\n"; - ofs << " \"pp_size\": " << metadata.parallel_config.pp_size << ",\n"; - ofs << " \"dp_size\": " << metadata.parallel_config.dp_size << ",\n"; - ofs << " \"sp_size\": " << metadata.parallel_config.sp_size << ",\n"; - ofs << " \"vpp_size\": " << metadata.parallel_config.vpp_size << "\n"; - ofs << " },\n"; + ofs << " \"tensors\": [\n"; for (size_t i = 0; i < metadata.tensors.size(); ++i) { const auto &tensor = metadata.tensors[i]; @@ -805,9 +802,9 @@ void Checkpoint::SaveMetadataFile(const std::filesystem::path &path, const Check } ofs << "],\n"; }; - write_segments("segment_global_offsets", &checkpoint::ShardSegment::global_offset); - write_segments("segment_local_offsets", &checkpoint::ShardSegment::local_offset); - write_segments("segment_lengths", &checkpoint::ShardSegment::length); + write_segments("segment_global_offsets", &ShardSegment::global_offset); + write_segments("segment_local_offsets", &ShardSegment::local_offset); + write_segments("segment_lengths", &ShardSegment::length); ofs << " \"file\": \"" << tensor.file << "\",\n"; ofs << " \"offset\": " << tensor.offset << ",\n"; ofs << " \"byte_size\": " << tensor.byte_size << ",\n"; diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index 34a01eb82..512eacd9f 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -6,22 +6,16 @@ #include #include #include -#include #include #include "glog/logging.h" #include "infini_train/include/checkpoint/checkpoint.h" -#include "infini_train/include/checkpoint/reshard.h" -#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" #include "infini_train/include/nn/parallel/ddp/distributed_optimizer.h" -#include "infini_train/include/nn/parallel/global.h" -#include "infini_train/include/nn/parallel/parallel_functional.h" -#include "infini_train/include/nn/parallel/work.h" -#include "infini_train/include/tensor.h" using namespace infini_train; namespace nn = infini_train::nn; @@ -29,7 +23,7 @@ namespace nn = infini_train::nn; namespace { std::filesystem::path ResolveCheckpointDirectory(const std::filesystem::path &root) { - const auto latest_path = root / "latest_checkpointed_iteration.txt"; + const auto latest_path = root / checkpoint::kLatestIterationFilename; if (!std::filesystem::exists(latest_path)) { return root; } @@ -41,52 +35,6 @@ std::filesystem::path ResolveCheckpointDirectory(const std::filesystem::path &ro return directory; } -void SynchronizeCheckpointRanks(const nn::Module &model) { - const auto parameters = model.Parameters(); - CHECK(!parameters.empty()) << "Cannot synchronize checkpoint save for a model without parameters"; - auto token = std::make_shared(std::vector{1}, DataType::kFLOAT32, parameters.front()->GetDevice()); - token->Fill(1.0f); - nn::parallel::function::AllReduce(token, nn::parallel::function::ReduceOpType::kSum, nullptr, true)->Synchronize(); -} - -void WaitForWriterManifests(const std::filesystem::path &staging_root, int tp_size, int pp_size, - int64_t expected_iteration) { - const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(10); - for (;;) { - bool ready = true; - for (int pp = 0; pp < pp_size && ready; ++pp) { - for (int tp = 0; tp < tp_size; ++tp) { - const int rank = nn::parallel::global::GetRankOf(0, tp, pp); - const auto manifest = staging_root / std::format("rank_{:06d}", rank) / "metadata.json"; - if (!std::filesystem::exists(manifest)) { - ready = false; - break; - } - const auto rank_metadata = Checkpoint::LoadMetadata(manifest.parent_path()); - if (!rank_metadata.has_metadata || rank_metadata.iteration != expected_iteration) { - ready = false; - break; - } - } - } - if (ready) { - return; - } - CHECK(std::chrono::steady_clock::now() < deadline) - << "Timed out waiting for checkpoint manifests in " << staging_root; - std::this_thread::sleep_for(std::chrono::milliseconds(10)); - } -} - -void WaitForGlobalMetadata(const std::filesystem::path &metadata_path) { - const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(10); - while (!std::filesystem::exists(metadata_path)) { - CHECK(std::chrono::steady_clock::now() < deadline) - << "Timed out waiting for global checkpoint metadata: " << metadata_path; - std::this_thread::sleep_for(std::chrono::milliseconds(10)); - } -} - } // namespace ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { @@ -98,25 +46,30 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & CHECK(dynamic_cast(args.optimizer.get()) == nullptr) << "Checkpoint restore does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; - // Resolve the checkpoint generation and load the global shard metadata. - auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); - CHECK(std::filesystem::exists(checkpoint_dir / "metadata.json")) + const auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); + CHECK(std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) << "Checkpoint metadata.json not found: " << checkpoint_dir; - auto metadata = Checkpoint::LoadMetadata(checkpoint_dir); - - CHECK(metadata.has_metadata); - CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; - // Reconstruct model and optimizer state for the current parallel topology. - checkpoint::LoadDistributedCheckpoint(checkpoint_dir, *args.model, args.optimizer.get(), args.state, - args.lr_scheduler.get(), metadata); - - // Validate architecture invariants before restoring training progress. - CHECK_EQ(args.state.n_layer, args.model_config.n_layer); - CHECK_EQ(args.state.n_head, args.model_config.n_head); - CHECK_EQ(args.state.n_kv_head, args.model_config.n_kv_head); - CHECK_EQ(args.state.n_embd, args.model_config.n_embd); - CHECK_GE(args.state.vocab_size, args.model_config.original_vocab_size) - << "Checkpoint vocabulary cannot represent the configured logical vocabulary"; + Checkpoint::Load(checkpoint_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get()); + + CHECK_EQ(args.state.n_layer, args.model_config.n_layer) + << "n_layer mismatch: ckpt=" << args.state.n_layer << ", config=" << args.model_config.n_layer; + CHECK_EQ(args.state.n_head, args.model_config.n_head) + << "n_head mismatch: ckpt=" << args.state.n_head << ", config=" << args.model_config.n_head; + CHECK_EQ(args.state.n_kv_head, args.model_config.n_kv_head) + << "n_kv_head mismatch: ckpt=" << args.state.n_kv_head << ", config=" << args.model_config.n_kv_head; + CHECK_EQ(args.state.n_embd, args.model_config.n_embd) + << "n_embd mismatch: ckpt=" << args.state.n_embd << ", config=" << args.model_config.n_embd; + if (args.state.original_vocab_size > 0) { + CHECK_EQ(args.state.original_vocab_size, args.model_config.original_vocab_size) + << "original_vocab_size mismatch: ckpt=" << args.state.original_vocab_size + << ", config=" << args.model_config.original_vocab_size; + CHECK_GE(args.state.padded_vocab_size, args.state.original_vocab_size) + << "Checkpoint padded vocabulary cannot represent its logical vocabulary"; + } else { + // Legacy trainer_state.json only stored the padded vocab size. + CHECK_GE(args.state.padded_vocab_size, args.model_config.original_vocab_size) + << "Legacy checkpoint vocabulary cannot represent the configured logical vocabulary"; + } result.global_step = static_cast(args.state.global_step); result.consumed_train_samples = static_cast(std::max(args.state.consumed_train_samples, 0)); if (args.rank.IsMainRank()) { @@ -129,6 +82,8 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & void SaveCheckpoint(const SaveCheckpointArgs &args) { CHECK(dynamic_cast(args.optimizer) == nullptr) << "Checkpoint save does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; + CHECK_GT(args.original_vocab_size, 0); + CHECK_GE(args.padded_vocab_size, args.original_vocab_size); const auto checkpoint_start = std::chrono::high_resolution_clock::now(); // Snapshot training progress and the topology that produced this checkpoint. TrainerState state{.global_step = args.global_step, @@ -137,7 +92,8 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { .n_head = args.n_head, .n_kv_head = args.n_kv_head, .n_embd = args.n_embd, - .vocab_size = args.vocab_size, + .original_vocab_size = args.original_vocab_size, + .padded_vocab_size = args.padded_vocab_size, .ddp_size = args.ddp_size, .tp_size = args.tp_size, .sp_size = args.sp_size, @@ -146,72 +102,11 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { const auto iteration_dir = args.checkpoint_root_dir.empty() ? args.save_dir : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); - std::filesystem::create_directories(iteration_dir); - - // Reset the manifest staging area before writer ranks publish their metadata. - const auto staging_root = iteration_dir / ".metadata_tmp"; - if (args.rank.IsMainRank()) { - std::filesystem::remove_all(staging_root); - } - SynchronizeCheckpointRanks(args.model); - - int dp_rank = 0, tp_rank = 0, pp_rank = 0; - nn::parallel::global::GetCoordOf(args.rank.GlobalRank(), dp_rank, tp_rank, pp_rank); - // DP ranks hold replicas; only one DP replica writes each TP/PP shard. - if (dp_rank != 0) { - return; - } - - const auto rank_dir = iteration_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); - std::filesystem::create_directories(rank_dir); - // Describe logical shards, plan their physical layout, and write this rank shard. - auto sharded_state = args.model.ShardedStateDict(); - std::unordered_map> optimizer_state; - if (args.optimizer != nullptr) { - optimizer_state = args.optimizer->StateDict(); - auto optimizer_sharded_state = checkpoint::BuildOptimizerShardedStateDict(sharded_state, optimizer_state); - sharded_state.Merge(std::move(optimizer_sharded_state)); - } - auto write_items = checkpoint::SavePlanner::Plan(sharded_state, args.rank.GlobalRank()); - Checkpoint::SaveSharded(rank_dir, sharded_state, write_items, args.model.StateDict(), optimizer_state, state, - args.rank.GlobalRank()); - - const auto staging_rank_dir = staging_root / std::format("rank_{:06d}", args.rank.GlobalRank()); - std::filesystem::create_directories(staging_rank_dir); - // Stage the local manifest until all writer ranks have completed their shards. - const auto local_manifest = staging_rank_dir / "metadata.json"; - if (std::filesystem::exists(local_manifest)) { - std::filesystem::remove(local_manifest); - } - std::filesystem::rename(rank_dir / "metadata.json", local_manifest); - - if (args.rank.IsMainRank()) { - Checkpoint::SaveTrainerStateFile(iteration_dir / "trainer_state.json", state); - if (args.lr_scheduler != nullptr) { - Checkpoint::SaveLRSchedulerStateFile(iteration_dir / "lr_scheduler.ckpt", args.lr_scheduler->StateDict()); - } - // Aggregate writer manifests and atomically publish the global metadata. - WaitForWriterManifests(staging_root, args.tp_size, args.pp_size, args.global_step); - auto global_metadata = Checkpoint::LoadMetadata(staging_root); - CHECK(global_metadata.has_metadata); - const auto temporary_metadata = iteration_dir / "metadata.json.tmp"; - const auto final_metadata = iteration_dir / "metadata.json"; - if (std::filesystem::exists(temporary_metadata)) { - std::filesystem::remove(temporary_metadata); - } - Checkpoint::SaveMetadataFile(temporary_metadata, global_metadata); - if (std::filesystem::exists(final_metadata)) { - std::filesystem::remove(final_metadata); - } - std::filesystem::rename(temporary_metadata, final_metadata); - std::filesystem::remove_all(staging_root); - } else { - WaitForGlobalMetadata(iteration_dir / "metadata.json"); - } + Checkpoint::Save(iteration_dir, args.model, args.optimizer, state, args.lr_scheduler); if (args.rank.IsMainRank() && !args.checkpoint_root_dir.empty()) { - const auto latest = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt"; - const auto temporary_latest = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt.tmp"; + const auto latest = args.checkpoint_root_dir / checkpoint::kLatestIterationFilename; + const auto temporary_latest = args.checkpoint_root_dir / checkpoint::kTemporaryLatestIterationFilename; { std::ofstream output(temporary_latest); CHECK(output.is_open()); @@ -230,6 +125,9 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { checkpoints.push_back(entry.path()); } } + // FIXME(jym): Pruning relies on lexicographic sorting of checkpoint directory names. + // This is only correct while iteration directories use zero-padded names (e.g. iter_0000042). + // If the naming convention changes to unpadded names, parse the iteration and sort numerically instead. std::sort(checkpoints.begin(), checkpoints.end()); while (checkpoints.size() > args.max_checkpoint_keep) { std::filesystem::remove_all(checkpoints.front()); diff --git a/infini_train/src/checkpoint/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc index 871070d19..708171955 100644 --- a/infini_train/src/checkpoint/load_planner.cc +++ b/infini_train/src/checkpoint/load_planner.cc @@ -41,21 +41,26 @@ int FragmentedAxis(const std::vector &axis_fragmentations) { return fragmented_axis; } -bool IsVocabularyTensor(const std::string &key) { - std::string parameter_key = key; - if (parameter_key.starts_with("adam.m.")) { - parameter_key = parameter_key.substr(7); - } else if (parameter_key.starts_with("adam.v.")) { - parameter_key = parameter_key.substr(7); +int ResolveShardDim(const std::string &key, const ShardedTensor &target, int saved_axis, + const std::vector &saved_global_shape) { + const int target_axis = FragmentedAxis(target.axis_fragmentations); + int shard_dim = target_axis >= 0 ? target_axis : saved_axis; + if (shard_dim < 0 && saved_global_shape != target.global_shape) { + shard_dim = 0; } - return parameter_key == "transformer.wte.weight" || parameter_key == "lm_head.weight"; + if (saved_axis >= 0 && target_axis >= 0) { + CHECK_EQ(saved_axis, target_axis) << "Shard dimension changed for tensor " << key; + } + return shard_dim; } -bool IsPaddingCompatible(const std::string &key, const std::vector &source, +bool IsPaddingCompatible(bool allow_shape_mismatch, const std::vector &source, const std::vector &target) { - if (!IsVocabularyTensor(key) || source.size() != target.size() || source.empty()) { + if (!allow_shape_mismatch || source.size() != target.size() || source.empty()) { return false; } + // Dim 0 is the padded vocabulary axis and may grow or shrink across TP layouts. + // Every non-padding dimension must still describe the same tensor shape. for (size_t dim = 1; dim < source.size(); ++dim) { if (source[dim] != target[dim]) { return false; @@ -106,7 +111,7 @@ LoadPlan LoadPlanner::PlanReshard(const Checkpoint::CheckpointMetadata &metadata ValidateCoordinates(key, source->global_shape, source->local_shape, source->global_offset, source->axis_fragmentations); CHECK(source->global_shape == target.global_shape - || IsPaddingCompatible(key, source->global_shape, target.global_shape)) + || IsPaddingCompatible(target.allow_shape_mismatch, source->global_shape, target.global_shape)) << "Global shape changed for tensor " << key; CHECK_EQ(FragmentedAxis(source->axis_fragmentations), saved_axis) << "Inconsistent saved shard dimensions for tensor " << key; @@ -165,15 +170,7 @@ LoadPlan LoadPlanner::PlanReshard(const Checkpoint::CheckpointMetadata &metadata continue; } - if (tensor_plan.shard_dim < 0) { - tensor_plan.shard_dim = saved_axis; - } - if (tensor_plan.shard_dim < 0 && candidates.front()->global_shape != target.global_shape) { - tensor_plan.shard_dim = 0; - } - if (saved_axis >= 0 && FragmentedAxis(target.axis_fragmentations) >= 0) { - CHECK_EQ(saved_axis, tensor_plan.shard_dim) << "Shard dimension changed for tensor " << key; - } + tensor_plan.shard_dim = ResolveShardDim(key, target, saved_axis, candidates.front()->global_shape); if (tensor_plan.shard_dim < 0) { const auto *source = candidates.front(); @@ -228,7 +225,7 @@ LoadPlan LoadPlanner::PlanReshard(const Checkpoint::CheckpointMetadata &metadata covered += read.length; } if (covered < target_length) { - CHECK(IsVocabularyTensor(key)) << "Incomplete target shard plan for " << key; + CHECK(target.allow_shape_mismatch) << "Incomplete target shard plan for " << key; CHECK_EQ(dim, 0) << "Vocabulary padding is only supported along dim 0"; CHECK_EQ(target_start + covered, candidates.front()->global_shape[0]) << "Only trailing vocabulary padding is supported for " << key; diff --git a/infini_train/src/checkpoint/reshard.cc b/infini_train/src/checkpoint/reshard.cc deleted file mode 100644 index fb2b364ca..000000000 --- a/infini_train/src/checkpoint/reshard.cc +++ /dev/null @@ -1,76 +0,0 @@ -#include "infini_train/include/checkpoint/reshard.h" - -#include -#include - -#include "glog/logging.h" - -#include "infini_train/include/checkpoint/load_planner.h" -#include "infini_train/include/checkpoint/load_strategy.h" -#include "infini_train/include/checkpoint/save_planner.h" -#include "infini_train/include/lr_scheduler.h" -#include "infini_train/include/nn/modules/module.h" -#include "infini_train/include/nn/parallel/global.h" -#include "infini_train/include/optimizer.h" - -namespace infini_train::checkpoint { - -void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, - TrainerState &state, LRScheduler *lr_scheduler, - const Checkpoint::CheckpointMetadata &metadata) { - CHECK(metadata.has_metadata); - CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; - // Build this rank's target shard layout and plan overlap reads from the saved source shards. - auto model_sharded_state = model.ShardedStateDict(); - auto plan = LoadPlanner::PlanReshard(metadata, model_sharded_state); - // Execute the read plan, assemble target tensors, and load the reconstructed model state. - IndexedRegionLoadStrategy strategy; - auto result = strategy.Execute(checkpoint_dir, plan); - model.LoadStateDict(result); - - // Restore training progress, but rewrite topology fields to describe the current runtime. - state = Checkpoint::LoadTrainerStateFile(checkpoint_dir / "trainer_state.json"); - const int current_tp = nn::parallel::global::GetTensorParallelSize(); - const int current_pp = nn::parallel::global::GetPipelineParallelSize(); - const bool topology_changed - = current_tp != metadata.parallel_config.tp_size || current_pp != metadata.parallel_config.pp_size; - state.tp_size = current_tp; - state.pp_size = current_pp; - state.ddp_size = nn::parallel::global::GetDataParallelSize(); - state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; - - // Reshard optimizer tensors only when TP or PP changed; otherwise load the matching writer shard directly. - if (optimizer != nullptr) { - if (topology_changed) { - const auto initialized_optimizer_state = optimizer->StateDict(); - auto optimizer_sharded_state - = BuildOptimizerShardedStateDict(model_sharded_state, initialized_optimizer_state); - auto optimizer_plan = LoadPlanner::PlanReshard(metadata, optimizer_sharded_state); - auto loaded_optimizer_state = strategy.Execute(checkpoint_dir, optimizer_plan); - optimizer->LoadStateDict(loaded_optimizer_state); - LOG(INFO) << "[CKPT] Resharded " << loaded_optimizer_state.size() - << " optimizer tensors across TP/PP topology change"; - } else { - int dp_rank = 0, tp_rank = 0, pp_rank = 0; - nn::parallel::global::GetCoordOf(nn::parallel::global::thread_global_rank, dp_rank, tp_rank, pp_rank); - const int writer_rank = nn::parallel::global::GetRankOf(0, tp_rank, pp_rank); - const auto optimizer_path = checkpoint_dir / std::format("rank_{:06d}/optimizer.ckpt", writer_rank); - CHECK(std::filesystem::exists(optimizer_path)) - << "Optimizer checkpoint not found for current_rank=" << nn::parallel::global::thread_global_rank - << ", coords=(dp=" << dp_rank << ", tp=" << tp_rank << ", pp=" << pp_rank - << "), writer_rank=" << writer_rank << ": " << optimizer_path; - LOG(INFO) << "[CKPT] Loading optimizer for current_rank=" << nn::parallel::global::thread_global_rank - << " from writer_rank=" << writer_rank << ": " << optimizer_path; - optimizer->LoadStateDict(Checkpoint::LoadStateDictFile(optimizer_path)); - } - } - // Scheduler state is topology-independent and can be restored directly. - if (lr_scheduler != nullptr && std::filesystem::exists(checkpoint_dir / "lr_scheduler.ckpt")) { - lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(checkpoint_dir / "lr_scheduler.ckpt")); - } - LOG(INFO) << "[CKPT] Restored " << result.size() - << " tensors with overlap reads from TP=" << metadata.parallel_config.tp_size - << ", PP=" << metadata.parallel_config.pp_size << " to TP=" << current_tp << ", PP=" << current_pp; -} - -} // namespace infini_train::checkpoint diff --git a/infini_train/src/checkpoint/save_planner.cc b/infini_train/src/checkpoint/save_planner.cc index 8b61b0b17..3a8004b28 100644 --- a/infini_train/src/checkpoint/save_planner.cc +++ b/infini_train/src/checkpoint/save_planner.cc @@ -2,6 +2,8 @@ #include "glog/logging.h" +#include "infini_train/include/checkpoint/constants.h" +#include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" namespace infini_train::checkpoint { @@ -11,24 +13,18 @@ BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, const std::unordered_map> &optimizer_state) { ShardedStateDict result; for (const auto &[key, tensor] : optimizer_state) { - if (key == "adam.t") { - ShardedTensor info; - info.key = key; + if (key == optimizers::kAdamStepKey) { + auto info = MakeShardedTensor(key, tensor->Dtype(), tensor->Dims()); info.local_key = key; - info.dtype = tensor->Dtype(); - info.global_shape = tensor->Dims(); - info.local_shape = tensor->Dims(); - info.global_offset.assign(tensor->Dims().size(), 0); - info.axis_fragmentations.assign(tensor->Dims().size(), 1); result.tensors.emplace(key, std::move(info)); continue; } std::string parameter_key; - if (key.starts_with("adam.m.")) { - parameter_key = key.substr(std::string("adam.m.").size()); - } else if (key.starts_with("adam.v.")) { - parameter_key = key.substr(std::string("adam.v.").size()); + if (key.starts_with(optimizers::kAdamFirstMomentPrefix)) { + parameter_key = key.substr(optimizers::kAdamFirstMomentPrefix.size()); + } else if (key.starts_with(optimizers::kAdamSecondMomentPrefix)) { + parameter_key = key.substr(optimizers::kAdamSecondMomentPrefix.size()); } else { CHECK(false) << "Unsupported optimizer state key: " << key; } @@ -48,17 +44,11 @@ BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, std::vector SavePlanner::Plan(const ShardedStateDict &sd, int rank) { std::vector items; - uint64_t model_offset = 0; - uint64_t optim_offset = 0; - for (auto &[key, info] : sd.tensors) { - bool is_optimizer = key.starts_with("adam."); - uint64_t &offset = is_optimizer ? optim_offset : model_offset; - + bool is_optimizer = key.starts_with(optimizers::kAdamOptimizerPrefix); WriteItem item; item.key = key; - item.filename = is_optimizer ? "optimizer.ckpt" : "model.ckpt"; - item.offset = offset; + item.filename = is_optimizer ? kOptimizerCheckpointFilename : kModelCheckpointFilename; item.byte_size = TensorByteSize(info.dtype, info.local_shape); item.dtype = info.dtype; item.local_shape = info.local_shape; @@ -67,7 +57,6 @@ std::vector SavePlanner::Plan(const ShardedStateDict &sd, int rank) { item.rank = rank; items.push_back(std::move(item)); - offset += items.back().byte_size; } return items; diff --git a/infini_train/src/core/ccl/ccl.cc b/infini_train/src/core/ccl/ccl.cc index 1bddee0e0..d0ece09c6 100644 --- a/infini_train/src/core/ccl/ccl.cc +++ b/infini_train/src/core/ccl/ccl.cc @@ -54,6 +54,11 @@ void CclImpl::ReduceScatter(const void *sendbuff, void *recvbuff, size_t recv_co LOG(FATAL) << "CclImpl::ReduceScatter is not implemented."; } +void CclImpl::AlltoAll(const void *sendbuff, void *recvbuff, size_t count, DataType dtype, const CclComm *comm, + Stream *stream) const { + LOG(FATAL) << "CclImpl::AlltoAll is not implemented."; +} + void CclImpl::Send(const void *buff, size_t count, DataType dtype, int peer, const CclComm *comm, Stream *stream) const { LOG(FATAL) << "CclImpl::Send is not implemented."; diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.cc b/infini_train/src/core/ccl/cuda/nccl_impl.cc index 9e4b1a0d6..5c622405b 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.cc +++ b/infini_train/src/core/ccl/cuda/nccl_impl.cc @@ -146,6 +146,48 @@ void NcclImpl::ReduceScatter(const void *sendbuff, void *recvbuff, size_t recv_c kNcclReduceOpMap.at(reduce_op), GetNcclComm(comm), GetCudaStream(stream))); } +void NcclImpl::AlltoAll(const void *sendbuff, void *recvbuff, size_t count, DataType dtype, const CclComm *comm, + Stream *stream) const { + auto nccl_comm = GetNcclComm(comm); + auto cuda_stream = GetCudaStream(stream); + CHECK_NE(sendbuff, recvbuff) << "NcclImpl::AlltoAll does not support in-place operation."; + + // NCCL 2.28.3+ provides native host collective ncclAlltoAll with the same contiguous rank-major layout. + // Older NCCL releases do not expose it, so fall back to an equivalent grouped ncclSend/ncclRecv schedule. +#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 28, 3) + NCCL_CHECK(ncclAlltoAll(sendbuff, recvbuff, count, kNcclDtypeMap.at(dtype), nccl_comm, cuda_stream)); +#else + int nranks = 0; + int rank = 0; + NCCL_CHECK(ncclCommCount(nccl_comm, &nranks)); + NCCL_CHECK(ncclCommUserRank(nccl_comm, &rank)); + CHECK_GT(nranks, 0); + CHECK_GE(rank, 0); + CHECK_LT(rank, nranks); + + const size_t chunk_bytes = count * kDataTypeToSize.at(dtype); + auto send_ptr = static_cast(sendbuff); + auto recv_ptr = static_cast(recvbuff); + + if (chunk_bytes > 0) { + CUDA_CHECK(cudaMemcpyAsync(recv_ptr + static_cast(rank) * chunk_bytes, + send_ptr + static_cast(rank) * chunk_bytes, chunk_bytes, + cudaMemcpyDeviceToDevice, cuda_stream)); + } + + NCCL_CHECK(ncclGroupStart()); + for (int peer = 0; peer < nranks; ++peer) { + if (peer == rank) { + continue; + } + const auto offset = static_cast(peer) * chunk_bytes; + NCCL_CHECK(ncclSend(send_ptr + offset, count, kNcclDtypeMap.at(dtype), peer, nccl_comm, cuda_stream)); + NCCL_CHECK(ncclRecv(recv_ptr + offset, count, kNcclDtypeMap.at(dtype), peer, nccl_comm, cuda_stream)); + } + NCCL_CHECK(ncclGroupEnd()); +#endif +} + void NcclImpl::Send(const void *buff, size_t count, DataType dtype, int peer, const CclComm *comm, Stream *stream) const { NCCL_CHECK(ncclSend(buff, count, kNcclDtypeMap.at(dtype), peer, GetNcclComm(comm), GetCudaStream(stream))); diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.h b/infini_train/src/core/ccl/cuda/nccl_impl.h index fca177fd9..42d0664a0 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.h +++ b/infini_train/src/core/ccl/cuda/nccl_impl.h @@ -42,6 +42,9 @@ class NcclImpl final : public CclImpl { nn::parallel::function::ReduceOpType reduce_op, const CclComm *comm, Stream *stream) const override; + void AlltoAll(const void *sendbuff, void *recvbuff, size_t count, DataType dtype, const CclComm *comm, + Stream *stream) const override; + void Send(const void *buff, size_t count, DataType dtype, int peer, const CclComm *comm, Stream *stream) const override; diff --git a/infini_train/src/nn/lora/lora_parallel_linear.cc b/infini_train/src/nn/lora/lora_parallel_linear.cc index 298b3e074..45130faf4 100644 --- a/infini_train/src/nn/lora/lora_parallel_linear.cc +++ b/infini_train/src/nn/lora/lora_parallel_linear.cc @@ -219,22 +219,17 @@ std::vector> LoRAColumnParallelLinear::LoRAParameters() return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } -checkpoint::ShardedStateDict LoRAColumnParallelLinear::ShardedStateDict(const std::string &prefix) const { - auto state = parallel::ColumnParallelLinear::ShardedStateDict(prefix); +ShardedStateDict LoRAColumnParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + auto state = parallel::ColumnParallelLinear::BuildShardedStateDict(prefix); const int tp_size = parallel::global::GetTensorParallelSize(); const auto &lora_a = parameter(kParamLoraAName); - checkpoint::ShardedTensor a; - a.key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; - a.dtype = lora_a->Dtype(); - a.global_shape = lora_a->Dims(); - a.local_shape = lora_a->Dims(); - a.global_offset = {0, 0}; - a.axis_fragmentations = {1, 1}; + const auto a_key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; + auto a = MakeShardedTensor(a_key, lora_a->Dtype(), lora_a->Dims()); state.tensors.emplace(a.key, std::move(a)); const auto &lora_b = parameter(kParamLoraBName); - checkpoint::ShardedTensor b; + ShardedTensor b; b.key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; b.dtype = lora_b->Dtype(); b.global_shape = {lora_b->Dims()[0] * tp_size, lora_b->Dims()[1]}; @@ -455,12 +450,12 @@ std::vector> LoRARowParallelLinear::LoRAParameters() con return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } -checkpoint::ShardedStateDict LoRARowParallelLinear::ShardedStateDict(const std::string &prefix) const { - auto state = parallel::RowParallelLinear::ShardedStateDict(prefix); +ShardedStateDict LoRARowParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + auto state = parallel::RowParallelLinear::BuildShardedStateDict(prefix); const int tp_size = parallel::global::GetTensorParallelSize(); const auto &lora_a = parameter(kParamLoraAName); - checkpoint::ShardedTensor a; + ShardedTensor a; a.key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; a.dtype = lora_a->Dtype(); a.global_shape = {lora_a->Dims()[0], lora_a->Dims()[1] * tp_size}; @@ -470,13 +465,8 @@ checkpoint::ShardedStateDict LoRARowParallelLinear::ShardedStateDict(const std:: state.tensors.emplace(a.key, std::move(a)); const auto &lora_b = parameter(kParamLoraBName); - checkpoint::ShardedTensor b; - b.key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; - b.dtype = lora_b->Dtype(); - b.global_shape = lora_b->Dims(); - b.local_shape = lora_b->Dims(); - b.global_offset = {0, 0}; - b.axis_fragmentations = {1, 1}; + const auto b_key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; + auto b = MakeShardedTensor(b_key, lora_b->Dtype(), lora_b->Dims()); state.tensors.emplace(b.key, std::move(b)); return state; } diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index ab8fd6ff8..7ad6f6624 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -188,29 +188,17 @@ std::unordered_map> Module::StateDict() con return state; } -checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix) const { - checkpoint::ShardedStateDict sd; +ShardedStateDict Module::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; for (auto &[name, param] : parameters_) { - checkpoint::ShardedTensor info; - info.key = prefix.empty() ? name : prefix + "." + name; - info.dtype = param->Dtype(); - info.global_shape = param->Dims(); - info.local_shape = param->Dims(); - info.global_offset.assign(param->Dims().size(), 0); - info.axis_fragmentations.assign(param->Dims().size(), 1); - sd.tensors[info.key] = std::move(info); + auto key = prefix.empty() ? name : prefix + "." + name; + sd.tensors.emplace(key, MakeShardedTensor(key, param->Dtype(), param->Dims())); } for (auto &[name, buffer] : buffers_) { - checkpoint::ShardedTensor info; - info.key = prefix.empty() ? name : prefix + "." + name; - info.dtype = buffer->Dtype(); - info.global_shape = buffer->Dims(); - info.local_shape = buffer->Dims(); - info.global_offset.assign(buffer->Dims().size(), 0); - info.axis_fragmentations.assign(buffer->Dims().size(), 1); - sd.tensors[info.key] = std::move(info); + auto key = prefix.empty() ? name : prefix + "." + name; + sd.tensors.emplace(key, MakeShardedTensor(key, buffer->Dtype(), buffer->Dims())); } for (auto &[name, module] : modules_) { @@ -219,7 +207,7 @@ checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix) } auto child_prefix = prefix.empty() ? name : prefix + "." + name; - auto child_sd = module->ShardedStateDict(child_prefix); + auto child_sd = module->BuildShardedStateDict(child_prefix); sd.Merge(std::move(child_sd)); } diff --git a/infini_train/src/nn/modules/normalization.cc b/infini_train/src/nn/modules/normalization.cc index 388b04de5..4bd0b0dab 100644 --- a/infini_train/src/nn/modules/normalization.cc +++ b/infini_train/src/nn/modules/normalization.cc @@ -7,6 +7,7 @@ #include "infini_train/include/device.h" #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/init.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" namespace infini_train::nn { @@ -18,6 +19,10 @@ LayerNorm::LayerNorm(const std::vector &normalized_shape, float eps, De = std::make_shared(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad(); parameters_[kParamBiasName] = std::make_shared(normalized_shape, DataType::kFLOAT32, device_)->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamWeightName]->set_sequence_parallel(true); + parameters_[kParamBiasName]->set_sequence_parallel(true); + } ResetParameters(); } @@ -35,6 +40,9 @@ void LayerNorm::ResetParameters() { RMSNorm::RMSNorm(int64_t dim, float eps, Device device) : CloneableModule(kType), eps_(eps) { parameters_[kParamWeightName] = std::make_shared(std::vector{dim}, DataType::kFLOAT32, device)->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamWeightName]->set_sequence_parallel(true); + } nn::init::Ones(parameters_[kParamWeightName]); } diff --git a/infini_train/src/nn/modules/transformer/causal_self_attention.cc b/infini_train/src/nn/modules/transformer/causal_self_attention.cc index 57a341fbc..0d2f48759 100644 --- a/infini_train/src/nn/modules/transformer/causal_self_attention.cc +++ b/infini_train/src/nn/modules/transformer/causal_self_attention.cc @@ -27,6 +27,10 @@ CausalSelfAttention::CausalSelfAttention(const TransformerConfig &config) : Clon if (config_.qk_layernorm) { modules_[kQNormLayerName] = std::make_shared(head_dim_, config_.norm_eps); modules_[kKNormLayerName] = std::make_shared(head_dim_, config_.norm_eps); + // NOTE(zbl): In Qwen3-8B, Q/K norm sees full sequences and TP-local heads. + // So we only need to finalize its replicated weights across TP regardless of SP. + modules_[kQNormLayerName]->parameter(RMSNorm::kParamWeightName)->set_sequence_parallel(false); + modules_[kKNormLayerName]->parameter(RMSNorm::kParamWeightName)->set_sequence_parallel(false); } int64_t qkv_dim = (config.n_head + 2 * n_kv_head_) * head_dim_; @@ -82,8 +86,8 @@ void CausalSelfAttention::SetupAttention(const TransformerConfig &config) { } } -checkpoint::ShardedStateDict CausalSelfAttention::ShardedStateDict(const std::string &prefix) const { - auto state = Module::ShardedStateDict(prefix); +ShardedStateDict CausalSelfAttention::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); const int tp_size = parallel::global::GetTensorParallelSize(); const int rank = parallel::tp_rank; const int64_t q_global = n_head_ * head_dim_; @@ -91,6 +95,8 @@ checkpoint::ShardedStateDict CausalSelfAttention::ShardedStateDict(const std::st const int64_t q_local = q_global / tp_size; const int64_t kv_local = kv_global / tp_size; + // FIXME(jym): Transformer should not depend on LoRA just to identify lora_B. Move the packed-QKV segment layout + // into a dedicated sharding abstraction shared by the base attention weight and LoRA parameters. const auto c_attn_prefix = prefix.empty() ? kCAttnLayerName : prefix + "." + kCAttnLayerName; auto set_qkv_segments = [&](const std::string ¶meter_name) { const auto key = c_attn_prefix + "." + parameter_name; diff --git a/infini_train/src/nn/modules/transformer/moe/router.cc b/infini_train/src/nn/modules/transformer/moe/router.cc index 252086846..b045400aa 100644 --- a/infini_train/src/nn/modules/transformer/moe/router.cc +++ b/infini_train/src/nn/modules/transformer/moe/router.cc @@ -11,6 +11,7 @@ #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/init.h" #include "infini_train/include/nn/modules/transformer/moe/moe_utils.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" namespace infini_train::nn::moe { @@ -24,12 +25,18 @@ TopKRouter::TopKRouter(const TransformerConfig &config) : CloneableModule(kType) = std::make_shared(std::vector{moe_config.num_experts, config_.n_embd}, DataType::kFLOAT32, device_) ->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamWeightName]->set_sequence_parallel(true); + } init::KaimingUniform(parameters_[kParamWeightName]); if (config_.add_bias_linear) { parameters_[kParamBiasName] = std::make_shared(std::vector{moe_config.num_experts}, DataType::kFLOAT32, device_) ->RequiresGrad(); + if (parallel::global::GetSequenceParallelEnabled()) { + parameters_[kParamBiasName]->set_sequence_parallel(true); + } parameters_[kParamBiasName]->Fill(0.0f); } } diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index 9704f9741..9411e9110 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -32,6 +32,9 @@ TransformerFirstStage::TransformerFirstStage(const TransformerConfig &config) // Only learned absolute position embedding uses a trainable WPE table. if (config_.position_embedding_type == PositionEmbeddingType::kLearnedAbsolute) { modules_[kWPELayerName] = std::make_shared(config_.block_size, config_.n_embd); + if (parallel::global::GetSequenceParallelEnabled()) { + modules_[kWPELayerName]->parameter(Embedding::kParamWeightName)->set_sequence_parallel(true); + } } else if (config_.position_embedding_type != PositionEmbeddingType::kRoPE) { LOG(FATAL) << "Unsupported position embedding type"; } @@ -234,17 +237,19 @@ TransformerModel::TransformerModel(const TransformerConfig config) auto chunk = std::make_shared(config_, start_layer, end_layer); start_layer_to_layer_size_and_chunk[start_layer] = std::make_pair(end_layer - start_layer, chunk); } - std::vector> h; + std::unordered_map> h; int chunk_idx = 0; for (auto &[start_layer, layer_size_and_chunk] : start_layer_to_layer_size_and_chunk) { auto [layer_size, chunk] = layer_size_and_chunk; - for (int idx = 0; idx < layer_size; ++idx) { - h.push_back(chunk->mutable_module(TransformerChunk::kHLayerName)->mutable_module(std::to_string(idx))); + for (int local_layer = 0; local_layer < layer_size; ++local_layer) { + const auto global_layer = start_layer + local_layer; + h[std::to_string(global_layer)] + = chunk->mutable_module(TransformerChunk::kHLayerName)->mutable_module(std::to_string(local_layer)); } modules_[kPPChunkNamePrefix + std::to_string(chunk_idx)] = std::move(chunk); ++chunk_idx; } - transformer[TransformerChunk::kHLayerName] = std::make_shared(std::move(h)); + transformer[TransformerChunk::kHLayerName] = std::make_shared(std::move(h)); } if (stage_info_.is_last_stage) { @@ -272,59 +277,15 @@ TransformerModel::TransformerModel(const TransformerConfig config) } } -namespace { - -std::vector GlobalLayerIndices(const parallel::StageInfo &stage_info) { - std::vector indices; - for (const auto &[start, end] : stage_info.layer_ranges_per_chunk) { - for (int layer = start; layer < end; ++layer) { indices.push_back(layer); } - } - std::sort(indices.begin(), indices.end()); - return indices; -} - -std::string RemapLayerKey(const std::string &key, const std::vector &from, const std::vector &to) { - const std::string marker - = std::string(TransformerModel::kTransformerModelName) + "." + TransformerChunk::kHLayerName + "."; - const auto marker_pos = key.find(marker); - if (marker_pos == std::string::npos) { - return key; - } - const auto index_start = marker_pos + marker.size(); - const auto index_end = key.find('.', index_start); - if (index_end == std::string::npos) { - return key; - } - int layer = -1; - try { - layer = std::stoi(key.substr(index_start, index_end - index_start)); - } catch (...) { return key; } - const auto it = std::find(from.begin(), from.end(), layer); - if (it == from.end()) { - return key; - } - const auto mapped = to[static_cast(std::distance(from.begin(), it))]; - return key.substr(0, index_start) + std::to_string(mapped) + key.substr(index_end); -} - -} // namespace - -checkpoint::ShardedStateDict TransformerModel::ShardedStateDict(const std::string &prefix) const { - auto local_state = Module::ShardedStateDict(prefix); - const auto global_layers = GlobalLayerIndices(stage_info_); - std::vector local_layers(global_layers.size()); - std::iota(local_layers.begin(), local_layers.end(), 0); - - checkpoint::ShardedStateDict global_state; - for (auto &[local_key, tensor] : local_state.tensors) { - const auto global_key = RemapLayerKey(local_key, local_layers, global_layers); - if (global_key != local_key) { - tensor.local_key = local_key; - tensor.key = global_key; - } - global_state.tensors.emplace(global_key, std::move(tensor)); +ShardedStateDict TransformerModel::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); + if (stage_info_.is_last_stage) { + const auto lm_head_prefix = prefix.empty() ? TransformerLastStage::kLMHeadLayerName + : prefix + "." + TransformerLastStage::kLMHeadLayerName; + const auto weight_key = lm_head_prefix + "." + parallel::ColumnParallelLinear::kParamWeightName; + state.tensors.at(weight_key).allow_shape_mismatch = true; } - return global_state; + return state; } std::vector>> @@ -335,18 +296,16 @@ TransformerModel::NamedParameters(const std::string &prefix, bool recurse, bool // Select public aliases so optimizer state keys match ShardedStateDict keys. auto parameters = Module::NamedParameters(prefix, true, false); - const auto sharded_state = ShardedStateDict(prefix); - const auto global_layers = GlobalLayerIndices(stage_info_); - std::vector local_layers(global_layers.size()); - std::iota(local_layers.begin(), local_layers.end(), 0); + const auto sharded_state = BuildShardedStateDict(prefix); std::vector>> result; std::unordered_set visited; + const auto private_pipeline_prefix = prefix.empty() ? "__pp" : prefix + ".__pp"; for (auto &[name, parameter] : parameters) { - name = RemapLayerKey(name, local_layers, global_layers); - if (!sharded_state.tensors.contains(name)) { + if (name.starts_with(private_pipeline_prefix)) { continue; } + CHECK(sharded_state.tensors.contains(name)) << "Parameter is missing from sharded state dict: " << name; if (remove_duplicate && !visited.insert(parameter.get()).second) { continue; } @@ -355,18 +314,6 @@ TransformerModel::NamedParameters(const std::string &prefix, bool recurse, bool return result; } -void TransformerModel::LoadStateDict(const std::unordered_map> &state_dict) { - const auto global_layers = GlobalLayerIndices(stage_info_); - std::vector local_layers(global_layers.size()); - std::iota(local_layers.begin(), local_layers.end(), 0); - - std::unordered_map> local_state; - for (const auto &[global_key, tensor] : state_dict) { - local_state.emplace(RemapLayerKey(global_key, global_layers, local_layers), tensor); - } - Module::LoadStateDict(local_state); -} - std::vector> TransformerModel::Forward(const std::vector> &x) { auto x1 = (*modules_[kPPFirstStageName])(x); for (int chunk_idx = 0; chunk_idx < stage_info_.layer_ranges_per_chunk.size(); ++chunk_idx) { diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index 100f5e5aa..8878944a4 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -220,8 +220,8 @@ std::unordered_map> DistributedDataParallel return modules_.at(kModuleName)->StateDict(); } -checkpoint::ShardedStateDict DistributedDataParallel::ShardedStateDict(const std::string &prefix) const { - return modules_.at(kModuleName)->ShardedStateDict(prefix); +ShardedStateDict DistributedDataParallel::BuildShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->BuildShardedStateDict(prefix); } void DistributedDataParallel::LoadStateDict( @@ -249,6 +249,10 @@ std::unique_ptr DistributedDataParallel::no_sync() { return std::make_unique([this, previous] { SetIsLastMicrobatch(previous); }); } +void DistributedDataParallel::FinishGradSync() { + for (auto &group : bucket_groups_) { group->FinishGradSync(); } +} + void DistributedDataParallel::SetIsLastMicrobatch(bool is_last_microbatch) { is_last_microbatch_->store(is_last_microbatch, std::memory_order_relaxed); if (reducer_) { diff --git a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc index 523bcf2d7..a8fa73e64 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_optimizer.cc @@ -68,6 +68,7 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShardParam &add size_t num_shard_params = 0; for (const auto &group : bucket_groups_) { + std::vector local_grad_shards; const bool use_grad_shard = group->config().zero_stage >= 2; const auto &buckets = group->buckets(); for (size_t bucket_idx = 0; bucket_idx < buckets.size(); ++bucket_idx) { @@ -107,6 +108,7 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShardParam &add auto param_piece = std::make_shared(*bucket_param, param_piece_offset_bytes, std::vector{static_cast(piece_numel)}); + param_piece->set_sequence_parallel(param->sequence_parallel()); auto grad_piece = std::make_shared(*bucket_grad, grad_piece_offset_bytes, std::vector{static_cast(piece_numel)}); @@ -115,24 +117,18 @@ void DistributedOptimizer::BuildShardParamsAndBindGrads(const AddShardParam &add // NOTE(zbl): Do not call `param->set_grad(grad_piece);` under ZeRO-2. // The base optimizer updates param_piece views only; original param->grad() // would be a partial flattened shard and does not represent the full parameter grad. + local_grad_shards.emplace_back(param, grad_piece); add_shard_param(param, param_piece); ++num_shard_params; } } + group->set_local_grad_shards(std::move(local_grad_shards)); } CHECK_GT(num_shard_params, 0) << "DistributedOptimizer: this DP rank owns no param pieces. " << "Check bucket padding/divisibility and param bucketing order."; } -void DistributedOptimizer::StartGradSync() { - for (auto &group : bucket_groups_) { group->StartGradSync(); } -} - -void DistributedOptimizer::FinishGradSync() { - for (auto &group : bucket_groups_) { group->FinishGradSync(); } -} - void DistributedOptimizer::StartParamSync(bool force_sync) { for (auto &group : bucket_groups_) { group->StartParamSync(force_sync); } } @@ -171,14 +167,10 @@ float DistributedOptimizer::learning_rate() const { } void DistributedOptimizer::Step() { - // 1. Ensure grads are synced - FinishGradSync(); - - // 2. Base optimizer step on owned param pieces CHECK(base_optimizer_) << "DistributedOptimizer: base optimizer is null."; base_optimizer_->Step(); - // 3. Gather updated param shards back to full params + // Gather updated param shards back to full params StartParamSync(/*force_sync=*/false); // TODO(zbl): Delay sync call until param is actually used in next step FinishParamSync(/*skip_next_bucket_dispatch=*/true); diff --git a/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc b/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc index 8a3c2d052..970c30f27 100644 --- a/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc +++ b/infini_train/src/nn/parallel/ddp/param_and_grad_buffer.cc @@ -240,6 +240,12 @@ std::shared_ptr ParamAndGradBucketGroup::GetLocalGradShardBuffer(size_t return grad_shard_buffer_list_[bucket_idx]; } +const std::vector &ParamAndGradBucketGroup::local_grad_shards() const { return local_grad_shards_; } + +void ParamAndGradBucketGroup::set_local_grad_shards(std::vector shards) { + local_grad_shards_ = std::move(shards); +} + void ParamAndGradBucketGroup::StartGradSync() { if (!collective_pg_) { LOG(FATAL) << "ParamAndGradBucketGroup: StartGradSync() called with null collective_pg_."; diff --git a/infini_train/src/nn/parallel/parallel_functional.cc b/infini_train/src/nn/parallel/parallel_functional.cc index ffd218d71..0ed4148f6 100644 --- a/infini_train/src/nn/parallel/parallel_functional.cc +++ b/infini_train/src/nn/parallel/parallel_functional.cc @@ -40,6 +40,15 @@ std::shared_ptr ReduceScatter(const std::shared_ptr &output, const return pg->ReduceScatter(output, input, reduce_op, async_op); } +std::shared_ptr AlltoAll(const std::shared_ptr &output, const std::shared_ptr &input, + const ProcessGroup *pg, bool async_op) { + auto device = output->GetDevice().type(); + if (pg == nullptr) { + pg = ProcessGroupFactory::Instance(device)->GetDefaultProcessGroup(); + } + return pg->AlltoAll(output, input, async_op); +} + std::vector>> Scatter(const std::vector> &input_tensors, const std::vector &devices, int dim) { std::vector>> output_tensors; diff --git a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc index 0dc34fdfb..f938c28fb 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc @@ -114,8 +114,8 @@ std::unordered_map> PipelineParallel::State return modules_.at(kModuleName)->StateDict(); } -checkpoint::ShardedStateDict PipelineParallel::ShardedStateDict(const std::string &prefix) const { - return modules_.at(kModuleName)->ShardedStateDict(prefix); +ShardedStateDict PipelineParallel::BuildShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->BuildShardedStateDict(prefix); } void PipelineParallel::LoadStateDict(const std::unordered_map> &state_dict) { diff --git a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc index 0ea2a4710..40ba16ddf 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -15,6 +15,7 @@ #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/nn/parallel/pp/pipeline_stage.h" #include "infini_train/include/nn/parallel/pp/send_recv.h" +#include "infini_train/include/nn/parallel/utils.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -299,6 +300,7 @@ float PipelineSchedule::Step(std::shared_ptr input, std::shared_ptrchunks()); optimizer->Step(); return lossf; diff --git a/infini_train/src/nn/parallel/process_group.cc b/infini_train/src/nn/parallel/process_group.cc index a557493e6..22e61b24a 100644 --- a/infini_train/src/nn/parallel/process_group.cc +++ b/infini_train/src/nn/parallel/process_group.cc @@ -313,6 +313,33 @@ std::shared_ptr ProcessGroup::Scatter(const std::vector ProcessGroup::AlltoAll(const std::shared_ptr &output, + const std::shared_ptr &input, bool async_op) const { + auto device = input->GetDevice(); + CHECK_EQ(device, output->GetDevice()); + CHECK(input->Dtype() == output->Dtype()); + CHECK_EQ(input->NumElements(), output->NumElements()); + CHECK_EQ(input->NumElements() % world_size_, 0) << "AlltoAll input must be evenly divisible by world size"; + core::DeviceGuard guard(device); + auto *compute_stream = runtime_impl_->GetStream(device); + auto *comm_stream = device_stream_map_.at(device.index()); + auto comm = device_comm_map_.at(device.index()); + + auto work = std::make_shared(device, comm); + runtime_impl_->EventRecord(work->ready_event(), compute_stream); + runtime_impl_->StreamWaitEvent(comm_stream, work->ready_event(), 0); + ccl_impl_->AlltoAll(input->DataPtr(), output->DataPtr(), input->NumElements() / world_size_, input->Dtype(), comm, + comm_stream); + runtime_impl_->EventRecord(work->done_event(), comm_stream); + + if (async_op) { + return work; + } else { + work->WaitNonBlocking(); + return nullptr; + } +} + std::shared_ptr ProcessGroup::Send(std::vector> tensors, int dest_rank, bool async_op) const { CHECK_GT(tensors.size(), 0); @@ -367,6 +394,46 @@ std::shared_ptr ProcessGroup::Recv(std::vector> te } } +std::shared_ptr ProcessGroup::BatchSendRecv(const std::vector &ops, bool async_op) const { + CHECK_GT(ops.size(), 0); + CHECK_NOTNULL(ops[0].tensor); + auto device = ops[0].tensor->GetDevice(); + core::DeviceGuard guard(device); + auto *compute_stream = runtime_impl_->GetStream(device); + auto *comm_stream = device_stream_map_.at(device.index()); + auto comm = device_comm_map_.at(device.index()); + + auto work = std::make_shared(device, comm); + runtime_impl_->EventRecord(work->ready_event(), compute_stream); + runtime_impl_->StreamWaitEvent(comm_stream, work->ready_event(), 0); + + { + core::CclGroupGuard ccl_group_guard(backend_); + for (const auto &op : ops) { + CHECK_NOTNULL(op.tensor); + CHECK_EQ(device, op.tensor->GetDevice()); + CHECK_GE(op.peer_rank, 0); + CHECK_LT(op.peer_rank, world_size_); + if (op.type == P2POpType::kSend) { + ccl_impl_->Send(op.tensor->DataPtr(), op.tensor->NumElements(), op.tensor->Dtype(), op.peer_rank, comm, + comm_stream); + } else { + ccl_impl_->Recv(op.tensor->DataPtr(), op.tensor->NumElements(), op.tensor->Dtype(), op.peer_rank, comm, + comm_stream); + } + } + } + + runtime_impl_->EventRecord(work->done_event(), comm_stream); + + if (async_op) { + return work; + } else { + work->WaitNonBlocking(); + return nullptr; + } +} + std::vector> ProcessGroup::BroadCast_(const std::vector> &input_tensors) const { std::vector> outputs; diff --git a/infini_train/src/nn/parallel/tensor_parallel.cc b/infini_train/src/nn/parallel/tensor_parallel.cc index 2755c15fa..890d80d0b 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -284,12 +284,12 @@ bool ColumnParallelLinear::input_is_parallel() const { return input_is_parallel_ bool ColumnParallelLinear::skip_bias_add() const { return skip_bias_add_; } bool ColumnParallelLinear::sequence_parallel() const { return sequence_parallel_; } -checkpoint::ShardedStateDict ColumnParallelLinear::ShardedStateDict(const std::string &prefix) const { - checkpoint::ShardedStateDict sd; +ShardedStateDict ColumnParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensor w; + ShardedTensor w; w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; w.dtype = weight->Dtype(); w.global_shape = {output_size_per_partition_ * tp_size, weight->Dims()[1]}; @@ -301,7 +301,7 @@ checkpoint::ShardedStateDict ColumnParallelLinear::ShardedStateDict(const std::s // Bias is also split along dim=0 if (bias_) { auto &bias = parameter(kParamBiasName); - checkpoint::ShardedTensor b; + ShardedTensor b; b.key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; b.dtype = bias->Dtype(); b.global_shape = {static_cast(output_size_per_partition_ * tp_size)}; @@ -335,6 +335,9 @@ RowParallelLinear::RowParallelLinear(int64_t in_features, int64_t out_features, if (bias) { parameters_[kParamBiasName] = std::make_shared(std::vector{out_features}, DataType::kFLOAT32, device_)->RequiresGrad(); + if (sequence_parallel_) { + parameters_[kParamBiasName]->set_sequence_parallel(true); + } } LinearResetParameters(parameters_[kParamWeightName], bias ? parameters_[kParamBiasName] : nullptr); @@ -369,12 +372,12 @@ bool RowParallelLinear::input_is_parallel() const { return input_is_parallel_; } bool RowParallelLinear::skip_bias_add() const { return skip_bias_add_; } bool RowParallelLinear::sequence_parallel() const { return sequence_parallel_; } -checkpoint::ShardedStateDict RowParallelLinear::ShardedStateDict(const std::string &prefix) const { - checkpoint::ShardedStateDict sd; +ShardedStateDict RowParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensor w; + ShardedTensor w; w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; w.dtype = weight->Dtype(); w.global_shape = {weight->Dims()[0], input_size_per_partition_ * tp_size}; @@ -386,13 +389,8 @@ checkpoint::ShardedStateDict RowParallelLinear::ShardedStateDict(const std::stri // Bias is NOT sharded in RowParallelLinear if (bias_) { auto &bias = parameter(kParamBiasName); - checkpoint::ShardedTensor b; - b.key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; - b.dtype = bias->Dtype(); - b.global_shape = bias->Dims(); - b.local_shape = bias->Dims(); - b.global_offset = {0}; - b.axis_fragmentations = {1}; + const auto key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; + auto b = MakeShardedTensor(key, bias->Dtype(), bias->Dims()); sd.tensors[b.key] = std::move(b); } @@ -456,18 +454,19 @@ VocabParallelEmbedding::Forward(const std::vector> &inpu return {output}; } -checkpoint::ShardedStateDict VocabParallelEmbedding::ShardedStateDict(const std::string &prefix) const { - checkpoint::ShardedStateDict sd; +ShardedStateDict VocabParallelEmbedding::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensor w; + ShardedTensor w; w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; w.dtype = weight->Dtype(); w.global_shape = {vocab_size_global_, embedding_dim_}; w.local_shape = weight->Dims(); w.global_offset = {vocab_start_index_, 0}; w.axis_fragmentations = {tp_size, 1}; + w.allow_shape_mismatch = true; sd.tensors[w.key] = std::move(w); return sd; diff --git a/infini_train/src/nn/parallel/utils.cc b/infini_train/src/nn/parallel/utils.cc index ee28a1694..9d3e871ec 100644 --- a/infini_train/src/nn/parallel/utils.cc +++ b/infini_train/src/nn/parallel/utils.cc @@ -1,13 +1,66 @@ #include "infini_train/include/nn/parallel/utils.h" +#include + #include "glog/logging.h" #include "infini_train/include/nn/functional.h" +#include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/ddp/distributed_data_parallel.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/nn/parallel/reduce_op_type.h" #include "infini_train/include/tensor.h" namespace infini_train::nn::parallel { +namespace { +const ProcessGroup *GetTensorParallelGroup(const Tensor &tensor) { + const int global_rank = tensor.GetDevice().Rank().GlobalRank(); + return ProcessGroupFactory::Instance(tensor.GetDevice().type()) + ->Get(GetTensorParallelProcessGroupName(global_rank)); +} + +void FinalizeSequenceParallelGradients(const std::vector ¶m_grads) { + // SP replicas see different sequence shards, so replicated parameter grads + // must be summed across TP before the optimizer consumes them. + if (!global::GetSequenceParallelEnabled() || global::GetTensorParallelSize() <= 1 || param_grads.empty()) { + return; + } + + const ProcessGroup *tp_group = nullptr; + for (const auto &[param, grad] : param_grads) { + if (!param || !param->sequence_parallel() || !grad) { + continue; + } + + if (tp_group == nullptr) { + tp_group = GetTensorParallelGroup(*param); + CHECK_NOTNULL(tp_group); + } + tp_group->AllReduce(grad, function::ReduceOpType::kSum, /*async_op=*/false); + } +} + +void FinalizeQKNormWeightGradients(const nn::Module &model_chunk, const std::vector ¶m_grads) { + // CausalSelfAttention Q/K norm weights (e.g. in Qwen3) are shared across heads and replicated across TP ranks. + // Each rank computes gradients for TP-local heads, so sum them across TP regardless of SP. + // These weights have sequence_parallel=false and are not reduced by FinalizeSequenceParallelGradients. + std::unordered_set qk_norm_params; + for (const auto &[name, param] : model_chunk.NamedParameters()) { + if (name.ends_with("q_norm.weight") || name.ends_with("k_norm.weight")) { + qk_norm_params.insert(param.get()); + } + } + + for (const auto &[param, grad] : param_grads) { + if (param && param->requires_grad() && grad && qk_norm_params.contains(param.get())) { + const auto *tp_group = GetTensorParallelGroup(*param); + CHECK_NOTNULL(tp_group); + tp_group->AllReduce(grad, function::ReduceOpType::kSum, /*async_op=*/false); + } + } +} +} // namespace std::string GetDataParallelProcessGroupName(int global_rank) { return "DP" + std::to_string(global::GetGroupId(global::DP, global_rank)); @@ -60,4 +113,42 @@ std::shared_ptr GatherTensorParallelShard(const std::shared_ptr return nn::function::Concat(rank_major_shards, dim)->Contiguous(); } +// Call once after all microbatch backwards, before gradient clipping and optimizer updates. +void FinalizeModelGrads(const std::vector> &model_chunks) { + // Finish DP gradient synchronization for all chunks before TP reduction. + for (const auto &model_chunk : model_chunks) { + if (auto ddp = std::dynamic_pointer_cast(model_chunk)) { + ddp->FinishGradSync(); + } + } + + if (global::GetTensorParallelSize() <= 1) { + // DP synchronization is complete; skip TP-specific finalization. + return; + } + + for (const auto &model_chunk : model_chunks) { + // Pair original parameters with the gradient views consumed by the optimizer. + std::vector param_grads; + auto ddp = std::dynamic_pointer_cast(model_chunk); + if (ddp && ddp->ddp_config().zero_stage >= 1) { + // ZeRO-1/2 use DP-local gradient shard views registered by DistOpt. + for (const auto &group : ddp->bucket_groups()) { + const auto &shards = group->local_grad_shards(); + param_grads.insert(param_grads.end(), shards.begin(), shards.end()); + } + } else { + // Ordinary optimizer / ZeRO-0 use full gradients from param.grad(). + for (const auto ¶m : model_chunk->Parameters()) { param_grads.emplace_back(param, param->grad()); } + } + // TP SUM combines local-token gradients for SP norms, row bias, position embeddings and routers. + FinalizeSequenceParallelGradients(param_grads); + + // TP SUM combines local-head gradients for replicated Q/K norm weights regardless of SP. + FinalizeQKNormWeightGradients(*model_chunk, param_grads); + } + + // NOTE(zbl): Extend this entry for PP tied embeddings, MoE shared parameters and loss normalization. +} + } // namespace infini_train::nn::parallel diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 39b999c77..666a0e520 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -135,29 +135,30 @@ std::unordered_map> Adam::StateDict() const std::unordered_map> state; for (size_t i = 0; i < m_.size(); ++i) { const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i]; - state.emplace("adam.m." + suffix, m_[i]); - state.emplace("adam.v." + suffix, v_[i]); + state.emplace(std::string(kAdamFirstMomentPrefix) + suffix, m_[i]); + state.emplace(std::string(kAdamSecondMomentPrefix) + suffix, v_[i]); } auto t_tensor = std::make_shared(std::vector{}, DataType::kINT64, Device()); *static_cast(t_tensor->DataPtr()) = t_; - state.emplace("adam.t", t_tensor); + state.emplace(std::string(kAdamStepKey), t_tensor); return state; } void Adam::LoadStateDict(const std::unordered_map> &state_dict) { for (size_t i = 0; i < m_.size(); ++i) { const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i]; - const auto m_key = "adam.m." + suffix; - const auto v_key = "adam.v." + suffix; + const auto m_key = std::string(kAdamFirstMomentPrefix) + suffix; + const auto v_key = std::string(kAdamSecondMomentPrefix) + suffix; CHECK(state_dict.contains(m_key)) << "Missing optimizer state: " << m_key; CHECK(state_dict.contains(v_key)) << "Missing optimizer state: " << v_key; m_[i]->CopyFrom(state_dict.at(m_key)); v_[i]->CopyFrom(state_dict.at(v_key)); } - CHECK(state_dict.contains("adam.t")) << "Missing optimizer state: adam.t"; - const Tensor t_cpu = state_dict.at("adam.t")->To(Device()); + const std::string step_key(kAdamStepKey); + CHECK(state_dict.contains(step_key)) << "Missing optimizer state: " << step_key; + const Tensor t_cpu = state_dict.at(step_key)->To(Device()); t_ = *static_cast(t_cpu.DataPtr()); } } // namespace optimizers diff --git a/infini_train/src/tensor.cc b/infini_train/src/tensor.cc index 4e61e221b..70d9591af 100644 --- a/infini_train/src/tensor.cc +++ b/infini_train/src/tensor.cc @@ -57,6 +57,7 @@ Tensor::Tensor(const Tensor &tensor, size_t offset, const std::vector & : buffer_(tensor.buffer_), offset_(tensor.offset_ + offset), dims_(dims), num_elements_(std::accumulate(dims.begin(), dims.end(), 1, std::multiplies())), dtype_(tensor.dtype_) { CHECK_LE(offset_ + kDataTypeToSize.at(dtype_) * num_elements_, buffer_->Size()); + sequence_parallel_ = tensor.sequence_parallel_; } Tensor::Tensor(const float *data, const std::vector &dims, DataType dtype, Device device) @@ -104,6 +105,10 @@ size_t Tensor::NumElements() const { return num_elements_; } DataType Tensor::Dtype() const { return dtype_; } +void Tensor::set_sequence_parallel(bool enabled) { sequence_parallel_ = enabled; } + +bool Tensor::sequence_parallel() const { return sequence_parallel_; } + std::shared_ptr Tensor::Detach() const { return std::make_shared(*this, 0, dims_); } void Tensor::Fill(Scalar value) { @@ -128,9 +133,11 @@ Eigen::Map> Tensor::Eig Tensor Tensor::To(Device device) { const auto buffer_device = buffer_->GetDevice(); if (device == buffer_device) { - auto new_tensor = Tensor(*this, offset_, dims_); + auto new_tensor = Tensor(*this, 0, dims_); + new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; if (grad_) { - new_tensor.grad_ = std::make_unique(*grad_.get(), grad_->offset_, grad_->dims_); + new_tensor.grad_ = std::make_unique(*grad_.get(), 0, grad_->dims_); } return new_tensor; } @@ -170,15 +177,18 @@ Tensor Tensor::To(Device device) { } new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; return new_tensor; } Tensor Tensor::To(DataType dtype) { if (dtype == dtype_) { - auto new_tensor = Tensor(*this, offset_, dims_); + auto new_tensor = Tensor(*this, 0, dims_); + new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; if (grad_) { - new_tensor.grad_ = std::make_unique(*grad_.get(), grad_->offset_, grad_->dims_); + new_tensor.grad_ = std::make_unique(*grad_.get(), 0, grad_->dims_); } return new_tensor; } @@ -194,6 +204,7 @@ Tensor Tensor::To(DataType dtype) { } new_tensor.requires_grad_ = requires_grad_; + new_tensor.sequence_parallel_ = sequence_parallel_; return new_tensor; } diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 000e80e71..7868ccf1d 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -1,13 +1,17 @@ #include +#include +#include +#include #include #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/checkpoint/load_planner.h" #include "infini_train/include/checkpoint/load_strategy.h" #include "infini_train/include/checkpoint/save_planner.h" -#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/shard_spec.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/optimizer.h" @@ -62,9 +66,9 @@ TEST(ModuleNamedParametersTest, SupportsTorchStyleArgumentsAndSharedParameterDed class CheckpointSerializationTest : public test::InfiniTrainTest {}; TEST(ShardedStateDictTest, RejectsDuplicateKeysWhenMerging) { - checkpoint::ShardedStateDict destination; + ShardedStateDict destination; destination.tensors["weight"] = {.key = "weight"}; - checkpoint::ShardedStateDict source; + ShardedStateDict source; source.tensors["weight"] = {.key = "weight"}; EXPECT_DEATH(destination.Merge(std::move(source)), "Duplicate sharded state-dict key: weight"); @@ -82,7 +86,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { p2->Fill(-1.5f); *model1->mutable_parameter("bias") = p2; - auto opt1 = std::make_shared(model1->Parameters(), 0.01); + auto opt1 = optimizers::Adam::CreateNamed(0.01)(model1->NamedParameters()); TrainerState saved{.global_step = 42, .consumed_train_samples = 100}; Checkpoint::Save(dir, *model1, opt1.get(), saved, nullptr); @@ -93,7 +97,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { auto q2 = std::make_shared(std::vector{4}, DataType::kFLOAT32, GetDevice()); q2->Fill(0.0f); *model2->mutable_parameter("bias") = q2; - auto opt2 = std::make_shared(model2->Parameters(), 0.01); + auto opt2 = optimizers::Adam::CreateNamed(0.01)(model2->NamedParameters()); TrainerState loaded; Checkpoint::Load(dir, *model2, opt2.get(), loaded, nullptr); @@ -115,7 +119,7 @@ TEST_P(CheckpointSerializationTest, DirectMetadataOffsetSupportsColumnSlices) { for (int row = 0; row < 4; ++row) { for (int column = 0; column < 4; ++column) { values[row * 4 + column] = row * 10.0f + column; } } - auto path = dir / "model.ckpt"; + auto path = dir / checkpoint::kModelCheckpointFilename; Checkpoint::SaveStateDictFile(path, {{"matrix", matrix}}); constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("matrix") - 1 + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); @@ -126,7 +130,7 @@ TEST_P(CheckpointSerializationTest, DirectMetadataOffsetSupportsColumnSlices) { .target_shape = {4, 2}, .shard_dim = 1, .reads = {{.key = "matrix", - .filename = "model.ckpt", + .filename = checkpoint::kModelCheckpointFilename, .dtype = DataType::kFLOAT32, .global_shape = {4, 4}, .byte_size = sizeof(float) * 16, @@ -160,23 +164,23 @@ TEST_P(CheckpointSerializationTest, ConvertsSavedBF16TensorToFP32Target) { source_data[2] = 3.0f; source_data[3] = 4.0f; auto source_bf16 = std::make_shared(source_fp32->To(DataType::kBFLOAT16)); - Checkpoint::SaveStateDictFile(dir / "model.ckpt", {{"weight", source_bf16}}); - constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("weight") - 1 - + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + Checkpoint::SaveStateDictFile(dir / checkpoint::kModelCheckpointFilename, {{"weight", source_bf16}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("weight") - 1 + sizeof(int8_t) + + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); checkpoint::LoadPlan plan; plan.tensors["weight"] = {.key = "weight", - .dtype = DataType::kFLOAT32, - .global_shape = {2, 2}, - .target_shape = {2, 2}, - .reads = {{.key = "weight", - .filename = "model.ckpt", - .dtype = DataType::kBFLOAT16, - .global_shape = {2, 2}, - .byte_size = source_bf16->SizeInBytes(), - .data_offset = data_offset, - .shard_dim = -1, - .source_shape = {2, 2}}}}; + .dtype = DataType::kFLOAT32, + .global_shape = {2, 2}, + .target_shape = {2, 2}, + .reads = {{.key = "weight", + .filename = checkpoint::kModelCheckpointFilename, + .dtype = DataType::kBFLOAT16, + .global_shape = {2, 2}, + .byte_size = source_bf16->SizeInBytes(), + .data_offset = data_offset, + .shard_dim = -1, + .source_shape = {2, 2}}}}; checkpoint::IndexedRegionLoadStrategy strategy; const auto loaded = strategy.Execute(dir, plan).at("weight"); @@ -195,7 +199,7 @@ TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { auto source = std::make_shared(std::vector{5, 2}, DataType::kFLOAT32, Device()); auto *source_data = static_cast(source->DataPtr()); for (int i = 0; i < 10; ++i) { source_data[i] = static_cast(i); } - Checkpoint::SaveStateDictFile(dir / "model.ckpt", {{"lm_head.weight", source}}); + Checkpoint::SaveStateDictFile(dir / checkpoint::kModelCheckpointFilename, {{"lm_head.weight", source}}); constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("lm_head.weight") - 1 + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); @@ -206,17 +210,18 @@ TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { .local_shape = {5, 2}, .global_offset = {0, 0}, .axis_fragmentations = {1, 1}, - .file = "model.ckpt", + .file = checkpoint::kModelCheckpointFilename, .offset = data_offset, .byte_size = sizeof(float) * 10}); - checkpoint::ShardedStateDict target; + ShardedStateDict target; target.tensors["lm_head.weight"] = {.key = "lm_head.weight", .dtype = DataType::kFLOAT32, .global_shape = {8, 2}, .local_shape = {4, 2}, .global_offset = {4, 0}, - .axis_fragmentations = {2, 1}}; + .axis_fragmentations = {2, 1}, + .allow_shape_mismatch = true}; const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); ASSERT_EQ(plan.tensors.at("lm_head.weight").trailing_zero_fill, 3); @@ -236,9 +241,7 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { std::filesystem::create_directories(dir); Checkpoint::CheckpointMetadata metadata; metadata.version = 3; - metadata.iteration = 17; metadata.has_metadata = true; - metadata.parallel_config = {.tp_size = 2, .pp_size = 2, .dp_size = 1, .sp_size = 1, .vpp_size = 2}; metadata.tensors.push_back({.key = "layer.0.weight", .dtype_str = "float32", .global_shape = {8, 4}, @@ -250,21 +253,23 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { .byte_size = 64, .stored_on_ranks = {0}, .pp_rank = 0}); - Checkpoint::SaveMetadataFile(dir / "metadata.json", metadata); + Checkpoint::SaveMetadataFile(dir / checkpoint::kMetadataFilename, metadata); + std::ifstream metadata_file(dir / checkpoint::kMetadataFilename); + const std::string metadata_json((std::istreambuf_iterator(metadata_file)), + std::istreambuf_iterator()); + EXPECT_EQ(metadata_json.find("\"iteration\""), std::string::npos); + EXPECT_EQ(metadata_json.find("\"parallel_config\""), std::string::npos); + EXPECT_EQ(metadata_json.find("\"model_config\""), std::string::npos); auto loaded = Checkpoint::LoadMetadata(dir); ASSERT_TRUE(loaded.has_metadata); - EXPECT_EQ(loaded.iteration, 17); - EXPECT_EQ(loaded.parallel_config.tp_size, 2); - EXPECT_EQ(loaded.parallel_config.pp_size, 2); - EXPECT_EQ(loaded.parallel_config.vpp_size, 2); ASSERT_EQ(loaded.tensors.size(), 1); EXPECT_EQ(loaded.tensors[0].file, "rank_000000/model.ckpt"); EXPECT_EQ(loaded.tensors[0].global_offset, std::vector({0, 0})); EXPECT_EQ(loaded.tensors[0].axis_fragmentations, std::vector({2, 1})); ASSERT_EQ(loaded.tensors[0].segments.size(), 1); EXPECT_EQ(loaded.tensors[0].segments[0], - (checkpoint::ShardSegment{.global_offset = 0, .local_offset = 0, .length = 4})); + (ShardSegment{.global_offset = 0, .local_offset = 0, .length = 4})); std::filesystem::remove_all(dir); } @@ -282,8 +287,8 @@ Checkpoint::CheckpointMetadata::TensorEntry MakeSavedShard(const std::string &ke .file = file}; } -checkpoint::ShardedStateDict MakeTarget(const std::string &key, int count, int index, int64_t global_size) { - checkpoint::ShardedStateDict target; +ShardedStateDict MakeTarget(const std::string &key, int count, int index, int64_t global_size) { + ShardedStateDict target; target.tensors[key] = {.key = key, .dtype = DataType::kFLOAT32, .global_shape = {global_size, 4}, @@ -295,7 +300,7 @@ checkpoint::ShardedStateDict MakeTarget(const std::string &key, int count, int i } // namespace TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { - checkpoint::ShardedStateDict model; + ShardedStateDict model; model.tensors["c_attn.weight"] = { .key = "c_attn.weight", .dtype = DataType::kFLOAT32, @@ -326,13 +331,14 @@ TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { EXPECT_EQ(m.segments, model.tensors.at("c_attn.weight").segments); EXPECT_EQ(m.local_key, "adam.m.c_attn.weight"); EXPECT_EQ(m.dtype, moment->Dtype()); + EXPECT_EQ(m.allow_shape_mismatch, model.tensors.at("c_attn.weight").allow_shape_mismatch); const auto &t = optimizer.tensors.at("adam.t"); EXPECT_TRUE(t.global_shape.empty()); EXPECT_TRUE(t.local_shape.empty()); } TEST(CheckpointOptimizerShardingTest, AdamMomentUsesOptimizerStateDtype) { - checkpoint::ShardedStateDict model; + ShardedStateDict model; model.tensors["weight"] = {.key = "weight", .dtype = DataType::kBFLOAT16, .global_shape = {4, 4}, @@ -402,7 +408,7 @@ TEST(CheckpointLoadPlannerTest, UsesExplicitGlobalOffsetsForUnevenShards) { .global_offset = {3, 0}, .axis_fragmentations = {2, 1}, .file = "rank_1/model.ckpt"}}; - checkpoint::ShardedStateDict target; + ShardedStateDict target; target.tensors["weight"] = {.key = "weight", .dtype = DataType::kFLOAT32, .global_shape = {8, 4}, @@ -433,7 +439,7 @@ TEST(CheckpointLoadPlannerTest, QkvSegmentsUseDimZeroWhenTpIsOne) { }; metadata.tensors.push_back(std::move(saved)); - checkpoint::ShardedStateDict target; + ShardedStateDict target; target.tensors["c_attn.weight"] = { .key = "c_attn.weight", .dtype = DataType::kFLOAT32, @@ -471,7 +477,7 @@ TEST(CheckpointLoadPlannerTest, QkvSegmentsPreserveTargetLocalLayoutAcrossTpChan metadata.tensors.push_back(std::move(shard)); } - checkpoint::ShardedStateDict target; + ShardedStateDict target; target.tensors["c_attn.weight"] = { .key = "c_attn.weight", .dtype = DataType::kFLOAT32, diff --git a/tests/checkpoint/test_lr_scheduler_state.cc b/tests/checkpoint/test_lr_scheduler_state.cc index a0e07dd93..40817b2d9 100644 --- a/tests/checkpoint/test_lr_scheduler_state.cc +++ b/tests/checkpoint/test_lr_scheduler_state.cc @@ -5,6 +5,7 @@ #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/optimizer.h" @@ -61,7 +62,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { 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")); + EXPECT_TRUE(std::filesystem::exists(dir / checkpoint::kLRSchedulerFilename)); auto model2 = MakeModel(GetDevice()); auto opt2 = std::make_shared(model2->Parameters(), kBaseLR); @@ -92,7 +93,7 @@ TEST_P(LRSchedulerCheckpointTest, SkipsLRSchedulerStateWhenSchedulerIsNull) { TrainerState saved{.global_step = 3}; Checkpoint::Save(dir, *model1, nullptr, saved, nullptr); - EXPECT_FALSE(std::filesystem::exists(dir / "lr_scheduler.ckpt")); + EXPECT_FALSE(std::filesystem::exists(dir / checkpoint::kLRSchedulerFilename)); std::filesystem::remove_all(dir); } diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index efcaf8b3f..67217c188 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -6,8 +6,10 @@ #include "infini_train/include/checkpoint/checkpoint.h" #include "infini_train/include/checkpoint/checkpoint_manager.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -26,7 +28,8 @@ TEST_P(TrainerStateTest, DefaultValues) { EXPECT_EQ(state.n_head, 0); EXPECT_EQ(state.n_kv_head, 0); EXPECT_EQ(state.n_embd, 0); - EXPECT_EQ(state.vocab_size, 0); + EXPECT_EQ(state.original_vocab_size, 0); + EXPECT_EQ(state.padded_vocab_size, 0); EXPECT_EQ(state.ddp_size, 1); EXPECT_EQ(state.tp_size, 1); EXPECT_EQ(state.sp_size, 1); @@ -48,9 +51,9 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { Checkpoint::Save(dir, *model, opt.get(), saved, nullptr); - EXPECT_TRUE(std::filesystem::exists(dir / "trainer_state.json")); + EXPECT_TRUE(std::filesystem::exists(dir / checkpoint::kTrainerStateFilename)); - std::ifstream ifs(dir / "trainer_state.json"); + std::ifstream ifs(dir / checkpoint::kTrainerStateFilename); std::string content((std::istreambuf_iterator(ifs)), std::istreambuf_iterator()); EXPECT_NE(content.find("\"global_step\""), std::string::npos); EXPECT_NE(content.find("\"consumed_train_samples\""), std::string::npos); @@ -69,12 +72,13 @@ TEST_P(TrainerStateTest, RoundTrip) { .n_head = 16, .n_kv_head = 8, .n_embd = 1024, - .vocab_size = 128256, + .original_vocab_size = 128000, + .padded_vocab_size = 128256, .ddp_size = 2, .tp_size = 1, .sp_size = 1, .pp_size = 2, - .vpp_size = 4, + .vpp_size = 1, }; auto model1 = std::make_shared(1, 3, true, GetDevice()); @@ -100,10 +104,14 @@ TEST_P(TrainerStateTest, RoundTrip) { EXPECT_EQ(loaded.n_head, 16); EXPECT_EQ(loaded.n_kv_head, 8); EXPECT_EQ(loaded.n_embd, 1024); - EXPECT_EQ(loaded.vocab_size, 128256); - EXPECT_EQ(loaded.ddp_size, 2); - EXPECT_EQ(loaded.pp_size, 2); - EXPECT_EQ(loaded.vpp_size, 4); + EXPECT_EQ(loaded.original_vocab_size, 128000); + EXPECT_EQ(loaded.padded_vocab_size, 128256); + EXPECT_EQ(loaded.ddp_size, nn::parallel::global::GetDataParallelSize()); + EXPECT_EQ(loaded.tp_size, nn::parallel::global::GetTensorParallelSize()); + EXPECT_EQ(loaded.pp_size, nn::parallel::global::GetPipelineParallelSize()); + EXPECT_EQ(loaded.sp_size, + nn::parallel::global::GetSequenceParallelEnabled() ? nn::parallel::global::GetTensorParallelSize() : 1); + EXPECT_EQ(loaded.vpp_size, 1); std::filesystem::remove_all(dir); } diff --git a/tests/lora/test_lora.cc b/tests/lora/test_lora.cc index 70d83a583..01e2725a2 100644 --- a/tests/lora/test_lora.cc +++ b/tests/lora/test_lora.cc @@ -88,7 +88,7 @@ TEST_P(LoRATest, ParallelLoRAShardedStateDictIncludesAdapterParameters) { 4, 6, /*bias=*/false, /*gather_output=*/false, /*input_is_parallel=*/false, /*skip_bias_add=*/false, /*sequence_parallel=*/false); auto column = std::make_shared(column_base, config, 4, 6); - const auto column_state = column->ShardedStateDict("column"); + const auto column_state = column->BuildShardedStateDict("column"); ASSERT_TRUE(column_state.tensors.contains("column.lora_A")); ASSERT_TRUE(column_state.tensors.contains("column.lora_B")); EXPECT_EQ(column_state.tensors.at("column.lora_A").axis_fragmentations, (std::vector{1, 1})); @@ -99,7 +99,7 @@ TEST_P(LoRATest, ParallelLoRAShardedStateDictIncludesAdapterParameters) { 4, 6, /*bias=*/false, /*reduce_output=*/true, /*input_is_parallel=*/true, /*skip_bias_add=*/false, /*sequence_parallel=*/false); auto row = std::make_shared(row_base, config, 4, 6); - const auto row_state = row->ShardedStateDict("row"); + const auto row_state = row->BuildShardedStateDict("row"); ASSERT_TRUE(row_state.tensors.contains("row.lora_A")); ASSERT_TRUE(row_state.tensors.contains("row.lora_B")); EXPECT_EQ(row_state.tensors.at("row.lora_A").axis_fragmentations, diff --git a/tests/tensor/test_tensor_copy.cc b/tests/tensor/test_tensor_copy.cc index 88d597b8e..84dcbc434 100644 --- a/tests/tensor/test_tensor_copy.cc +++ b/tests/tensor/test_tensor_copy.cc @@ -27,4 +27,34 @@ TEST_P(TensorCopyTest, CopiesPreservesDataType) { EXPECT_EQ(target->Dtype(), DataType::kFLOAT32); } +TEST_P(TensorCopyTest, NoOpConversionsPreserveParameterMetadataAndViews) { + auto storage = std::make_shared(std::vector{12}, DataType::kFLOAT32, GetDevice()); + auto grad_storage = std::make_shared(std::vector{12}, DataType::kFLOAT32, GetDevice()); + auto parameter = std::make_shared(*storage, 4 * sizeof(float), std::vector{4}); + auto grad = std::make_shared(*grad_storage, 4 * sizeof(float), std::vector{4}); + parameter->RequiresGrad(); + parameter->set_sequence_parallel(true); + parameter->set_grad(grad); + + for (auto converted : {parameter->To(GetDevice()), parameter->To(parameter->Dtype())}) { + EXPECT_TRUE(converted.requires_grad()); + EXPECT_TRUE(converted.sequence_parallel()); + EXPECT_EQ(converted.DataPtr(), parameter->DataPtr()); + EXPECT_EQ(converted.Dims(), parameter->Dims()); + ASSERT_NE(converted.grad(), nullptr); + EXPECT_EQ(converted.grad()->DataPtr(), grad->DataPtr()); + EXPECT_EQ(converted.grad()->Dims(), grad->Dims()); + } +} + +TEST_P(TensorCopyTest, NoOpConversionsPreserveFrozenTensorMetadata) { + auto tensor = std::make_shared(std::vector{4}, DataType::kFLOAT32, GetDevice()); + for (auto converted : {tensor->To(GetDevice()), tensor->To(tensor->Dtype())}) { + EXPECT_FALSE(converted.requires_grad()); + EXPECT_FALSE(converted.sequence_parallel()); + EXPECT_EQ(converted.grad(), nullptr); + EXPECT_EQ(converted.DataPtr(), tensor->DataPtr()); + } +} + INFINI_TRAIN_REGISTER_TEST(TensorCopyTest); diff --git a/tests/transformer/test_transformer_architecture.cc b/tests/transformer/test_transformer_architecture.cc index c754ea725..f6ed318f6 100644 --- a/tests/transformer/test_transformer_architecture.cc +++ b/tests/transformer/test_transformer_architecture.cc @@ -163,7 +163,7 @@ TEST_P(TransformerModuleTest, LLaMA3Model) { auto model = std::make_shared(config); model->To(GetDevice()); EXPECT_FALSE(model->Parameters().empty()); - const auto sharded_state = model->ShardedStateDict(); + const auto sharded_state = model->BuildShardedStateDict(); for (const auto &[name, parameter] : model->NamedParameters()) { EXPECT_TRUE(sharded_state.tensors.contains(name)) << "Missing shard metadata for named parameter: " << name; } From ec23a11c869751818f47e93cfff2def19236c48d Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Fri, 9 Oct 2026 13:35:57 +0000 Subject: [PATCH 5/5] fix: complete checkpoint resharding follow-ups --- .../include/nn/modules/transformer/mlp.h | 2 + infini_train/src/checkpoint/checkpoint.cc | 159 ++++++---------- .../src/nn/modules/transformer/mlp.cc | 35 ++++ .../src/nn/modules/transformer/transformer.cc | 4 + .../test_checkpoint_serialization.cc | 180 +++++++++++++++++- tests/checkpoint/test_trainer_state.cc | 4 +- 6 files changed, 278 insertions(+), 106 deletions(-) diff --git a/infini_train/include/nn/modules/transformer/mlp.h b/infini_train/include/nn/modules/transformer/mlp.h index ecf5672b6..0248f51ad 100644 --- a/infini_train/include/nn/modules/transformer/mlp.h +++ b/infini_train/include/nn/modules/transformer/mlp.h @@ -19,5 +19,7 @@ class MLP : public infini_train::nn::CloneableModule { std::vector> Forward(const std::vector> &x) override; + + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; }; } // namespace infini_train::nn diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index a291ec698..c02784def 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -1,6 +1,8 @@ #include "infini_train/include/checkpoint/checkpoint.h" #include +#include +#include #include #include #include @@ -256,10 +258,6 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m CHECK(metadata.has_metadata); CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; state = LoadTrainerState(checkpoint_dir / checkpoint::kTrainerStateFilename); - // TODO(jym): Support VPP checkpoint resharding by describing virtual pipeline chunks in the target shard layout. - CHECK_EQ(state.vpp_size, 1) << "Checkpoint resharding with saved VPP is not supported yet"; - CHECK_EQ(nn::parallel::global::GetVirtualPipelineParallelSize(), 1) - << "Checkpoint resharding with runtime VPP is not supported yet"; auto model_sharded_state = model.BuildShardedStateDict(); checkpoint::IndexedRegionLoadStrategy strategy; @@ -272,6 +270,7 @@ void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &m state.pp_size = current_pp; state.ddp_size = nn::parallel::global::GetDataParallelSize(); state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; + state.vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); if (optimizer != nullptr) { auto optimizer_sharded_state @@ -416,11 +415,9 @@ Checkpoint::LoadStateDictFile(const std::filesystem::path &path) { } static std::string DataTypeToString(DataType dt) { - auto it = kDataTypeToDesc.find(dt); - if (it != kDataTypeToDesc.end()) { - return it->second; - } - return "fp32"; + const auto it = kDataTypeToDesc.find(dt); + CHECK(it != kDataTypeToDesc.end()) << "Unsupported checkpoint tensor dtype: " << static_cast(dt); + return it->second; } void Checkpoint::SaveLocalShard(const std::filesystem::path &checkpoint_dir, const ShardedStateDict &sharded_sd, @@ -564,6 +561,52 @@ static std::string ExtractJsonString(const std::string &obj, const std::string & return obj.substr(q1 + 1, q2 - q1 - 1); } +static std::string_view TrimWhitespace(std::string_view value) { + while (!value.empty() && std::isspace(static_cast(value.front()))) { value.remove_prefix(1); } + while (!value.empty() && std::isspace(static_cast(value.back()))) { value.remove_suffix(1); } + return value; +} + +template std::vector ExtractIntegerArray(const std::string &obj, const char *field_name) { + const auto field = std::string("\"") + field_name + "\""; + const auto field_pos = obj.find(field); + if (field_pos == std::string::npos) { + return {}; + } + + const auto array_start = obj.find('[', field_pos + field.size()); + const auto array_end = obj.find(']', array_start); + CHECK(array_start != std::string::npos && array_end != std::string::npos) + << "Invalid integer array in checkpoint metadata field " << field_name; + + std::string_view contents(obj.data() + array_start + 1, array_end - array_start - 1); + contents = TrimWhitespace(contents); + if (contents.empty()) { + return {}; + } + + std::vector values; + while (!contents.empty()) { + const auto delimiter = contents.find(','); + const auto token = TrimWhitespace(contents.substr(0, delimiter)); + CHECK(!token.empty()) << "Empty integer in checkpoint metadata field " << field_name; + + T value{}; + const auto [end, error] = std::from_chars(token.data(), token.data() + token.size(), value); + CHECK(error == std::errc{} && end == token.data() + token.size()) + << "Invalid integer in checkpoint metadata field " << field_name << ": " << token; + values.push_back(value); + + if (delimiter == std::string_view::npos) { + break; + } + contents.remove_prefix(delimiter + 1); + contents = TrimWhitespace(contents); + CHECK(!contents.empty()) << "Trailing comma in checkpoint metadata field " << field_name; + } + return values; +} + static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem::path &checkpoint_dir) { Checkpoint::CheckpointMetadata meta; auto metadata_path = checkpoint_dir / checkpoint::kMetadataFilename; @@ -633,85 +676,14 @@ static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem:: entry.byte_size = ExtractNumberField(obj, "byte_size", 0); entry.pp_rank = ExtractNumberField(obj, "pp_rank", 0); - // global_shape: [x, y, z] - auto gs_pos = obj.find("\"global_shape\""); - if (gs_pos != std::string::npos) { - auto b1 = obj.find('[', gs_pos); - auto b2 = obj.find(']', b1); - if (b1 != std::string::npos && b2 != std::string::npos) { - std::string gs = obj.substr(b1 + 1, b2 - b1 - 1); - std::stringstream ss(gs); - std::string tok; - while (std::getline(ss, tok, ',')) { - try { - entry.global_shape.push_back(std::stoll(tok)); - } catch (...) {} - } - } - } - - auto ls_pos = obj.find("\"local_shape\""); - if (ls_pos != std::string::npos) { - auto b1 = obj.find('[', ls_pos); - auto b2 = obj.find(']', b1); - std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); - std::string tok; - while (std::getline(ss, tok, ',')) { - try { - entry.local_shape.push_back(std::stoll(tok)); - } catch (...) {} - } - } + entry.global_shape = ExtractIntegerArray(obj, "global_shape"); + entry.local_shape = ExtractIntegerArray(obj, "local_shape"); + entry.global_offset = ExtractIntegerArray(obj, "global_offset"); + entry.axis_fragmentations = ExtractIntegerArray(obj, "axis_fragmentations"); - auto offset_pos = obj.find("\"global_offset\""); - if (offset_pos != std::string::npos) { - auto b1 = obj.find('[', offset_pos); - auto b2 = obj.find(']', b1); - std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); - std::string token; - while (std::getline(ss, token, ',')) { - try { - entry.global_offset.push_back(std::stoll(token)); - } catch (...) {} - } - } - - auto fragments_pos = obj.find("\"axis_fragmentations\""); - if (fragments_pos != std::string::npos) { - auto b1 = obj.find('[', fragments_pos); - auto b2 = obj.find(']', b1); - std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); - std::string token; - while (std::getline(ss, token, ',')) { - try { - entry.axis_fragmentations.push_back(std::stoi(token)); - } catch (...) {} - } - } - - auto extract_int64_array = [&](const char *name) { - std::vector values; - auto field_pos = obj.find(std::string("\"") + name + "\""); - if (field_pos == std::string::npos) { - return values; - } - auto b1 = obj.find('[', field_pos); - auto b2 = obj.find(']', b1); - if (b1 == std::string::npos || b2 == std::string::npos) { - return values; - } - std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); - std::string token; - while (std::getline(ss, token, ',')) { - try { - values.push_back(std::stoll(token)); - } catch (...) {} - } - return values; - }; - const auto segment_global_offsets = extract_int64_array("segment_global_offsets"); - const auto segment_local_offsets = extract_int64_array("segment_local_offsets"); - const auto segment_lengths = extract_int64_array("segment_lengths"); + const auto segment_global_offsets = ExtractIntegerArray(obj, "segment_global_offsets"); + const auto segment_local_offsets = ExtractIntegerArray(obj, "segment_local_offsets"); + const auto segment_lengths = ExtractIntegerArray(obj, "segment_lengths"); CHECK_EQ(segment_global_offsets.size(), segment_local_offsets.size()); CHECK_EQ(segment_global_offsets.size(), segment_lengths.size()); for (size_t i = 0; i < segment_lengths.size(); ++i) { @@ -720,18 +692,7 @@ static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem:: .length = segment_lengths[i]}); } - auto ranks_pos = obj.find("\"stored_on_ranks\""); - if (ranks_pos != std::string::npos) { - auto b1 = obj.find('[', ranks_pos); - auto b2 = obj.find(']', b1); - std::stringstream ss(obj.substr(b1 + 1, b2 - b1 - 1)); - std::string tok; - while (std::getline(ss, tok, ',')) { - try { - entry.stored_on_ranks.push_back(std::stoi(tok)); - } catch (...) {} - } - } + entry.stored_on_ranks = ExtractIntegerArray(obj, "stored_on_ranks"); meta.tensors.push_back(std::move(entry)); obj_pos = obj_end + 1; diff --git a/infini_train/src/nn/modules/transformer/mlp.cc b/infini_train/src/nn/modules/transformer/mlp.cc index 3cf94f91c..8c12c0cfa 100644 --- a/infini_train/src/nn/modules/transformer/mlp.cc +++ b/infini_train/src/nn/modules/transformer/mlp.cc @@ -93,4 +93,39 @@ MLP::Forward(const std::vector> &x) { return (*modules_[kCProjLayerName])(x2); } +ShardedStateDict MLP::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); + if (!modules_.contains(kSwiGLULayerName)) { + return state; + } + + const auto c_fc_prefix = prefix.empty() ? kCFcLayerName : prefix + "." + kCFcLayerName; + const auto weight_key = c_fc_prefix + "." + parallel::ColumnParallelLinear::kParamWeightName; + const auto weight_it = state.tensors.find(weight_key); + CHECK(weight_it != state.tensors.end()) << "Missing packed SwiGLU weight metadata: " << weight_key; + + const auto &weight = weight_it->second; + CHECK(!weight.global_shape.empty() && !weight.local_shape.empty()); + CHECK_EQ(weight.global_shape[0] % 2, 0); + CHECK_EQ(weight.local_shape[0] % 2, 0); + const int64_t global_rows_per_projection = weight.global_shape[0] / 2; + const int64_t local_rows_per_projection = weight.local_shape[0] / 2; + const int rank = parallel::tp_rank; + + for (auto &[key, tensor] : state.tensors) { + if (!key.starts_with(c_fc_prefix + ".") || tensor.global_shape.empty() || tensor.local_shape.empty() + || tensor.global_shape[0] != weight.global_shape[0] || tensor.local_shape[0] != weight.local_shape[0]) { + continue; + } + tensor.global_offset.assign(tensor.global_shape.size(), 0); + tensor.segments = { + {.global_offset = rank * local_rows_per_projection, .local_offset = 0, .length = local_rows_per_projection}, + {.global_offset = global_rows_per_projection + rank * local_rows_per_projection, + .local_offset = local_rows_per_projection, + .length = local_rows_per_projection}, + }; + } + return state; +} + } // namespace infini_train::nn diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index 9411e9110..d39a13483 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -267,6 +267,10 @@ TransformerModel::TransformerModel(const TransformerConfig config) // applied after loading weights so it won't be overwritten. Also fix GPT2::FromLLMC() loading logic to respect // weight tying (do not create/load a separate lm_head.weight tensor; load once into the tied weight) so // parameter counting matches PyTorch/PEFT. + // FIXME(jym): A checkpoint saved with PP > 1 contains independent wte.weight and lm_head.weight tensors. Loading + // it with PP == 1 aliases both keys to the same destination Tensor, so Module::LoadStateDict copies both values + // and the final result depends on unordered-map iteration order. Canonicalize tied checkpoint keys, or reject + // PP > 1 to PP == 1 resharding for tied models until cross-stage weight tying is supported. if (config_.tie_weights && nn::parallel::global::GetPipelineParallelSize() == 1) { // https://paperswithcode.com/method/weight-tying *mutable_module(kTransformerModelName) diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 7868ccf1d..59d6ab3bd 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -1,3 +1,4 @@ +#include #include #include #include @@ -11,10 +12,11 @@ #include "infini_train/include/checkpoint/load_planner.h" #include "infini_train/include/checkpoint/load_strategy.h" #include "infini_train/include/checkpoint/save_planner.h" -#include "infini_train/include/shard_spec.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/modules/transformer/mlp.h" #include "infini_train/include/optimizer.h" +#include "infini_train/include/shard_spec.h" #include "infini_train/include/tensor.h" #include "tests/common/test_utils.h" @@ -74,6 +76,24 @@ TEST(ShardedStateDictTest, RejectsDuplicateKeysWhenMerging) { EXPECT_DEATH(destination.Merge(std::move(source)), "Duplicate sharded state-dict key: weight"); } +TEST(ShardedStateDictTest, PackedSwiGLUUsesGateAndUpSegments) { + nn::TransformerConfig config; + config.n_embd = 4; + config.activation_type = nn::MLPType::kSwiGLU; + config.ffn_expansion_ratio = 3.0f; + config.ffn_dim_multiplier = std::nullopt; + config.multiple_of = 1; + + const auto state = nn::MLP(config).BuildShardedStateDict("mlp"); + const auto &weight = state.tensors.at("mlp.c_fc.weight"); + ASSERT_EQ(weight.segments.size(), 2); + EXPECT_EQ(weight.segments[0], (ShardSegment{.global_offset = 0, .local_offset = 0, .length = 8})); + EXPECT_EQ(weight.segments[1], (ShardSegment{.global_offset = 8, .local_offset = 8, .length = 8})); + + const auto &bias = state.tensors.at("mlp.c_fc.bias"); + EXPECT_EQ(bias.segments, weight.segments); +} + TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { auto dir = std::filesystem::temp_directory_path() / "test_ckpt_fp32"; std::filesystem::remove_all(dir); @@ -235,6 +255,29 @@ TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { std::filesystem::remove_all(dir); } +TEST(CheckpointMetadataTest, RejectsMalformedIntegerArrays) { + const auto dir = std::filesystem::temp_directory_path() / "test_malformed_checkpoint_metadata"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + { + std::ofstream metadata_file(dir / checkpoint::kMetadataFilename); + metadata_file << R"({ + "version": 3, + "format": "infinitrain_sharded", + "tensors": [ + { + "key": "weight", + "dtype": "float32", + "global_shape": [8, invalid] + } + ] +})"; + } + + EXPECT_DEATH(Checkpoint::LoadMetadata(dir), "Invalid integer in checkpoint metadata field global_shape"); + std::filesystem::remove_all(dir); +} + TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { auto dir = std::filesystem::temp_directory_path() / "test_global_metadata"; std::filesystem::remove_all(dir); @@ -255,8 +298,7 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { .pp_rank = 0}); Checkpoint::SaveMetadataFile(dir / checkpoint::kMetadataFilename, metadata); std::ifstream metadata_file(dir / checkpoint::kMetadataFilename); - const std::string metadata_json((std::istreambuf_iterator(metadata_file)), - std::istreambuf_iterator()); + const std::string metadata_json((std::istreambuf_iterator(metadata_file)), std::istreambuf_iterator()); EXPECT_EQ(metadata_json.find("\"iteration\""), std::string::npos); EXPECT_EQ(metadata_json.find("\"parallel_config\""), std::string::npos); EXPECT_EQ(metadata_json.find("\"model_config\""), std::string::npos); @@ -268,8 +310,7 @@ TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { EXPECT_EQ(loaded.tensors[0].global_offset, std::vector({0, 0})); EXPECT_EQ(loaded.tensors[0].axis_fragmentations, std::vector({2, 1})); ASSERT_EQ(loaded.tensors[0].segments.size(), 1); - EXPECT_EQ(loaded.tensors[0].segments[0], - (ShardSegment{.global_offset = 0, .local_offset = 0, .length = 4})); + EXPECT_EQ(loaded.tensors[0].segments[0], (ShardSegment{.global_offset = 0, .local_offset = 0, .length = 4})); std::filesystem::remove_all(dir); } @@ -297,6 +338,18 @@ ShardedStateDict MakeTarget(const std::string &key, int count, int index, int64_ .axis_fragmentations = {count, 1}}; return target; } + +std::shared_ptr MakeColumnTensor(const std::vector &values) { + auto tensor = std::make_shared(std::vector{static_cast(values.size()), 1}, + DataType::kFLOAT32, Device()); + std::copy(values.begin(), values.end(), static_cast(tensor->DataPtr())); + return tensor; +} + +uint64_t SingleTensorDataOffset(const std::string &key, size_t dimensions) { + return sizeof(uint32_t) * 3 + sizeof(uint32_t) + key.size() + sizeof(int8_t) + sizeof(uint32_t) + + sizeof(int64_t) * dimensions + sizeof(uint64_t); +} } // namespace TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { @@ -506,6 +559,106 @@ TEST(CheckpointLoadPlannerTest, QkvSegmentsPreserveTargetLocalLayoutAcrossTpChan } } +TEST(CheckpointLoadPlannerTest, PackedSwiGLUExecutesTensorParallelTwoToOne) { + const auto dir = std::filesystem::temp_directory_path() / "test_swiglu_tp2_to_tp1"; + std::filesystem::remove_all(dir); + const std::string key = "c_fc.weight"; + const auto data_offset = SingleTensorDataOffset(key, 2); + + Checkpoint::CheckpointMetadata metadata; + for (int rank = 0; rank < 2; ++rank) { + const auto rank_dir = dir / ("rank_" + std::to_string(rank)); + std::filesystem::create_directories(rank_dir); + const std::vector values = rank == 0 ? std::vector{0, 1, 4, 5} : std::vector{2, 3, 6, 7}; + const auto tensor = MakeColumnTensor(values); + Checkpoint::SaveStateDictFile(rank_dir / checkpoint::kModelCheckpointFilename, {{key, tensor}}); + + metadata.tensors.push_back({ + .key = key, + .dtype_str = "fp32", + .global_shape = {8, 1}, + .local_shape = {4, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = { + {.global_offset = rank * 2, .local_offset = 0, .length = 2}, + {.global_offset = 4 + rank * 2, .local_offset = 2, .length = 2}, + }, + .file = "rank_" + std::to_string(rank) + "/" + checkpoint::kModelCheckpointFilename, + .offset = data_offset, + .byte_size = tensor->SizeInBytes(), + }); + } + + ShardedStateDict target; + target.tensors[key] = { + .key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {8, 1}, + .local_shape = {8, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 4}, + {.global_offset = 4, .local_offset = 4, .length = 4}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto loaded = checkpoint::IndexedRegionLoadStrategy().Execute(dir, plan).at(key); + const auto *data = static_cast(loaded->DataPtr()); + for (int row = 0; row < 8; ++row) { EXPECT_FLOAT_EQ(data[row], static_cast(row)); } + std::filesystem::remove_all(dir); +} + +TEST(CheckpointLoadPlannerTest, PackedSwiGLUExecutesTensorParallelOneToTwo) { + const auto dir = std::filesystem::temp_directory_path() / "test_swiglu_tp1_to_tp2"; + const auto rank_dir = dir / "rank_0"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(rank_dir); + const std::string key = "c_fc.weight"; + const auto tensor = MakeColumnTensor({0, 1, 2, 3, 4, 5, 6, 7}); + Checkpoint::SaveStateDictFile(rank_dir / checkpoint::kModelCheckpointFilename, {{key, tensor}}); + + Checkpoint::CheckpointMetadata metadata; + metadata.tensors.push_back({ + .key = key, + .dtype_str = "fp32", + .global_shape = {8, 1}, + .local_shape = {8, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 4}, + {.global_offset = 4, .local_offset = 4, .length = 4}, + }, + .file = "rank_0/" + std::string(checkpoint::kModelCheckpointFilename), + .offset = SingleTensorDataOffset(key, 2), + .byte_size = tensor->SizeInBytes(), + }); + + ShardedStateDict target; + target.tensors[key] = { + .key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {8, 1}, + .local_shape = {4, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = { + {.global_offset = 2, .local_offset = 0, .length = 2}, + {.global_offset = 6, .local_offset = 2, .length = 2}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto loaded = checkpoint::IndexedRegionLoadStrategy().Execute(dir, plan).at(key); + const auto *data = static_cast(loaded->DataPtr()); + const std::vector expected = {2, 3, 6, 7}; + for (size_t row = 0; row < expected.size(); ++row) { EXPECT_FLOAT_EQ(data[row], expected[row]); } + std::filesystem::remove_all(dir); +} + TEST(CheckpointLoadPlannerTest, PipelineReshardPlansOnlyTargetStageKeys) { Checkpoint::CheckpointMetadata metadata; metadata.tensors = {MakeSavedShard("layer.0.weight", 1, 0, 8, "old_pp0/model.ckpt"), @@ -516,3 +669,20 @@ TEST(CheckpointLoadPlannerTest, PipelineReshardPlansOnlyTargetStageKeys) { ASSERT_EQ(plan.tensors.at("layer.1.weight").reads.size(), 1); EXPECT_EQ(plan.tensors.at("layer.1.weight").reads[0].filename, "old_pp1/model.ckpt"); } + +TEST(CheckpointLoadPlannerTest, VirtualPipelineReshardUsesGlobalLayerKeys) { + Checkpoint::CheckpointMetadata metadata; + for (int layer = 0; layer < 4; ++layer) { + const auto key = "transformer.h." + std::to_string(layer) + ".weight"; + metadata.tensors.push_back( + MakeSavedShard(key, 1, 0, 8, "saved_layer_" + std::to_string(layer) + "/model.ckpt")); + } + + auto target = MakeTarget("transformer.h.0.weight", 1, 0, 8); + target.Merge(MakeTarget("transformer.h.2.weight", 1, 0, 8)); + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + + ASSERT_EQ(plan.tensors.size(), 2); + EXPECT_EQ(plan.tensors.at("transformer.h.0.weight").reads[0].filename, "saved_layer_0/model.ckpt"); + EXPECT_EQ(plan.tensors.at("transformer.h.2.weight").reads[0].filename, "saved_layer_2/model.ckpt"); +} diff --git a/tests/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index 67217c188..6df573892 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -78,7 +78,7 @@ TEST_P(TrainerStateTest, RoundTrip) { .tp_size = 1, .sp_size = 1, .pp_size = 2, - .vpp_size = 1, + .vpp_size = 2, }; auto model1 = std::make_shared(1, 3, true, GetDevice()); @@ -111,7 +111,7 @@ TEST_P(TrainerStateTest, RoundTrip) { EXPECT_EQ(loaded.pp_size, nn::parallel::global::GetPipelineParallelSize()); EXPECT_EQ(loaded.sp_size, nn::parallel::global::GetSequenceParallelEnabled() ? nn::parallel::global::GetTensorParallelSize() : 1); - EXPECT_EQ(loaded.vpp_size, 1); + EXPECT_EQ(loaded.vpp_size, nn::parallel::global::GetVirtualPipelineParallelSize()); std::filesystem::remove_all(dir); }