From d5d42c3ca63a6bf5c2e363d94b38aaad4b4b77fc Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Fri, 10 Jul 2026 18:27:14 +0800 Subject: [PATCH 1/3] feat: support checkpoint resharding --- infini_train/include/checkpoint/checkpoint.h | 52 +++ .../include/checkpoint/load_planner.h | 49 +++ infini_train/include/checkpoint/reshard.h | 96 +++++ .../include/checkpoint/save_planner.h | 69 ++++ infini_train/include/checkpoint/shard_spec.h | 51 +++ infini_train/include/nn/modules/module.h | 6 + .../include/nn/parallel/tensor_parallel.h | 13 + infini_train/src/checkpoint/checkpoint.cc | 272 +++++++++++++ .../src/checkpoint/checkpoint_manager.cc | 287 +++++++++++-- infini_train/src/checkpoint/load_planner.cc | 146 +++++++ infini_train/src/checkpoint/reshard.cc | 379 ++++++++++++++++++ infini_train/src/checkpoint/save_planner.cc | 32 ++ infini_train/src/nn/modules/module.cc | 39 ++ .../src/nn/parallel/tensor_parallel.cc | 113 +++++- 14 files changed, 1569 insertions(+), 35 deletions(-) create mode 100644 infini_train/include/checkpoint/load_planner.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/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 bf784d51..8d493479 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,6 +42,53 @@ class Checkpoint { static void Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, bool load_optimizer_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 Optimizer *optimizer, const TrainerState &state, bool save_optimizer_state, + const LRScheduler *lr_scheduler, 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::string file; + uint64_t offset = 0; + uint64_t byte_size = 0; + std::vector stored_on_ranks; + }; + + std::vector tensors; + bool has_metadata = false; + }; + + static CheckpointMetadata LoadMetadata(const std::filesystem::path &checkpoint_dir); + + + static void SaveLRSchedulerStateFile(const std::filesystem::path &path, const LRSchedulerStateDict &state_dict); + static LRSchedulerStateDict LoadLRSchedulerStateFile(const std::filesystem::path &path); + + + 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); diff --git a/infini_train/include/checkpoint/load_planner.h b/infini_train/include/checkpoint/load_planner.h new file mode 100644 index 00000000..bcaa5982 --- /dev/null +++ b/infini_train/include/checkpoint/load_planner.h @@ -0,0 +1,49 @@ +#pragma once + +#include +#include +#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" +#include "infini_train/include/tensor.h" + +namespace infini_train::checkpoint { + +struct ReadItem { + std::string key; + std::string filename; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector shard_specs; + int replica_id = 0; + uint64_t byte_size = 0; + + bool IsLocalRead(int current_tp_rank, int current_tp_size, int current_pp_rank, int current_pp_size) const; +}; + + +class LoadPlanner { +public: + struct PlanResult { + std::vector read_items; + std::unordered_map>> loaded_sd; + }; + + + static PlanResult PlanAndLoad(const std::filesystem::path &checkpoint_dir, + const Checkpoint::CheckpointMetadata &metadata); + + + static std::vector Plan(const Checkpoint::CheckpointMetadata &metadata); + + + static std::unordered_map> + LoadFile(const std::filesystem::path &checkpoint_dir, const std::string &filename); +}; + +} // 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 00000000..c457b4ad --- /dev/null +++ b/infini_train/include/checkpoint/reshard.h @@ -0,0 +1,96 @@ +#pragma once + +#include +#include +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/tensor.h" + +namespace infini_train { +class Optimizer; +class LRScheduler; +namespace nn { +class Module; +} +} // namespace infini_train + +namespace infini_train::checkpoint { + + +struct ReshardSlice { + int dim; + int64_t global_offset; + int64_t size; +}; + + +struct ReshardTensorPlan { + std::string key; + std::vector saved_local_shape; + std::vector target_shape; + std::vector slices; + bool needs_reshard = false; + bool is_replicated = false; +}; + + +struct ReshardPlan { + std::vector tensors; + int saved_tp_size = 1; + int saved_pp_size = 1; + int current_tp_size = 1; + int current_pp_size = 1; + bool tp_changed = false; + bool pp_changed = false; +}; + + + +class ReshardPlanner { +public: + + static ReshardPlan ComputePlan(const std::vector &read_items, + const Checkpoint::CheckpointMetadata::ParallelConfig &saved_config, + const Checkpoint::CheckpointMetadata::ParallelConfig ¤t_config, int tp_rank, + int pp_rank); + + + enum class PartitionType { + kColumnParallel, + kRowParallel, + kReplicated, + }; + + static PartitionType InferPartitionType(const std::string &key); +}; + + +class ReshardExecutor { +public: + + + + + static std::unordered_map> + Execute(const ReshardPlan &plan, const std::unordered_map> &full_sd, + int tp_rank, int pp_rank); +}; + + + + +std::unordered_map> +AllGatherFullModel(const std::unordered_map> &local_sd, int saved_tp_size); + + +void ReshardAndLoad(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, + const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank); + +} // 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 00000000..192ec0f1 --- /dev/null +++ b/infini_train/include/checkpoint/save_planner.h @@ -0,0 +1,69 @@ +#pragma once + +#include +#include +#include + +#include "infini_train/include/checkpoint/shard_spec.h" +#include "infini_train/include/datatype.h" + +namespace infini_train::checkpoint { + + +struct WriteItem { + std::string key; + std::string filename; // "model.ckpt" or "optimizer.ckpt" + uint64_t offset = 0; + uint64_t byte_size = 0; + DataType dtype = DataType::kFLOAT32; + std::vector local_shape; + std::vector shard_specs; + int replica_id = 0; + int rank = 0; +}; + + +class SavePlanner { +public: + static std::vector Plan(const ShardedStateDict &sd, int rank); +}; + + +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; + } +} + + +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 00000000..0cb16d11 --- /dev/null +++ b/infini_train/include/checkpoint/shard_spec.h @@ -0,0 +1,51 @@ +#pragma once + +#include +#include +#include +#include + +#include "infini_train/include/datatype.h" + +namespace infini_train::checkpoint { + + +struct ShardSpec { + int dim = -1; + int shard_count = 1; + int shard_index = 0; + std::string parallel_type; // "tp" / "pp" / "dp" / "sp" + + bool operator==(const ShardSpec &other) const { + return dim == other.dim && shard_count == other.shard_count && shard_index == other.shard_index + && parallel_type == other.parallel_type; + } + bool operator!=(const ShardSpec &other) const { return !(*this == other); } +}; + + +struct ShardedTensorInfo { + std::string key; // "transformer.h.0.attn.c_attn.weight" + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector local_shape; + std::vector shard_specs; + int replica_id = 0; + + bool operator==(const ShardedTensorInfo &other) const { + return key == other.key && dtype == other.dtype && global_shape == other.global_shape + && local_shape == other.local_shape && shard_specs == other.shard_specs; + } +}; + + +struct ShardedStateDict { + std::map tensors; + + + void Merge(ShardedStateDict &&other) { + for (auto &[key, info] : other.tensors) { tensors.emplace(std::move(key), std::move(info)); } + } +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index ca32b4d9..070dcb8d 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,6 +64,11 @@ class Module : public std::enable_shared_from_this { std::unordered_map> StateDict() const; + + virtual checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", + const std::vector &parent_shards + = {}) 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); diff --git a/infini_train/include/nn/parallel/tensor_parallel.h b/infini_train/include/nn/parallel/tensor_parallel.h index 0611dfbb..c7a17429 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,10 @@ class ColumnParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", + const std::vector &parent_shards + = {}) const override; + protected: bool bias_ = true; bool gather_output_ = false; // whether to return full local output tensor after forward (need gather) @@ -66,6 +71,10 @@ class RowParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", + const std::vector &parent_shards + = {}) const override; + protected: bool bias_ = true; bool reduce_output_ = false; // whether to return full local output tensor after forward (need reduce) @@ -85,6 +94,10 @@ class VocabParallelEmbedding : public nn::CloneableModule> Forward(const std::vector> &input_tensors) override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", + const std::vector &parent_shards + = {}) 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 34eb5a0d..c997c948 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -11,6 +11,7 @@ #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/optimizer.h" @@ -350,4 +351,275 @@ 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::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); +} + +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); +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- + +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 Optimizer *optimizer, const TrainerState &state, bool save_optimizer_state, + const LRScheduler *lr_scheduler, 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; + + + { + std::unordered_map> filtered_sd; + for (const auto &[key, info] : sharded_sd.tensors) { + if (info.replica_id != 0) { + continue; + } + + if (key.starts_with("adam.")) { + continue; + } + + auto it = state_dict.find(key); + if (it != state_dict.end()) { + filtered_sd.emplace(key, it->second); + } + } + if (!filtered_sd.empty()) { + SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); + } + } + + + if (save_optimizer_state && optimizer != nullptr) { + auto opt_state = optimizer->StateDict(); + if (!opt_state.empty()) { + SaveStateDict(checkpoint_dir / "optimizer.ckpt", opt_state); + } + } + + + if (lr_scheduler != nullptr) { + SaveLRSchedulerState(checkpoint_dir / "lr_scheduler.ckpt", lr_scheduler->StateDict()); + } + + + SaveTrainerState(checkpoint_dir / "trainer_state.json", state); + + + { + std::ofstream ofs(checkpoint_dir / "metadata.json"); + CHECK(ofs.is_open()) << "Failed to open metadata.json: " << checkpoint_dir / "metadata.json"; + + ofs << "{\n"; + ofs << " \"version\": 1,\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"; + + for (size_t i = 0; i < write_items.size(); ++i) { + auto &item = write_items[i]; + if (item.replica_id != 0) { + continue; + } + + auto it = sharded_sd.tensors.find(item.key); + if (it == sharded_sd.tensors.end()) { + continue; + } + + 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"; + + // shard_specs + if (!it->second.shard_specs.empty()) { + ofs << " \"shard_specs\": [\n"; + for (size_t s = 0; s < it->second.shard_specs.size(); ++s) { + auto &spec = it->second.shard_specs[s]; + ofs << " {\"dim\": " << spec.dim << ", \"shard_count\": " << spec.shard_count + << ", \"shard_index\": " << spec.shard_index << ", \"parallel_type\": \"" << spec.parallel_type + << "\"}"; + if (s + 1 < it->second.shard_specs.size()) { + ofs << ","; + } + ofs << "\n"; + } + ofs << " ],\n"; + } + + ofs << " \"file\": \"" << item.filename << "\",\n"; + ofs << " \"offset\": " << item.offset << ",\n"; + ofs << " \"byte_size\": " << item.byte_size << ",\n"; + ofs << " \"stored_on_ranks\": [" << global_rank << "]\n"; + ofs << " }"; + if (i + 1 < write_items.size()) { + ofs << ","; + } + ofs << "\n"; + } + + ofs << " ]\n"; + ofs << "}\n"; + + LOG(INFO) << "[CKPT] metadata.json written"; + } + + LOG(ERROR) << "[CKPT] SaveSharded done: dir=" << checkpoint_dir; +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- + +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); +} + +Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::path &checkpoint_dir) { + 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); + + + 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); + + + size_t obj_pos = 0; + while ((obj_pos = tensor_block.find('{', obj_pos)) != std::string::npos) { + size_t obj_end = tensor_block.find('}', obj_pos); + 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); + + // 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 (...) {} + } + } + } + + 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; +} } // namespace infini_train diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index d3cbb750..fd843b18 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -1,24 +1,186 @@ #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/load_planner.h" +#include "infini_train/include/checkpoint/reshard.h" +#include "infini_train/include/checkpoint/save_planner.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/process_group.h" +#include "infini_train/include/nn/parallel/rank.h" +#include "infini_train/include/nn/parallel/tensor_parallel.h" +#include "infini_train/include/nn/parallel/utils.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. +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- +static std::unordered_map> LoadRankShards(const std::filesystem::path &ckpt_dir, + int tp_world_size) { + std::unordered_map> merged_sd; + std::unordered_set all_keys; + + + for (int r = 0; r < tp_world_size; ++r) { + auto rank_dir = ckpt_dir / std::format("rank_{:06d}", r); + auto model_path = rank_dir / "model.ckpt"; + if (!std::filesystem::exists(model_path)) { + + model_path = ckpt_dir / "model.ckpt"; + } + if (!std::filesystem::exists(model_path)) { + LOG(WARNING) << "[CKPT] No model.ckpt found for rank " << r; + continue; + } + auto sd = Checkpoint::LoadStateDictFile(model_path); + for (const auto &[key, tensor] : sd) { all_keys.insert(key); } + } + + + auto device = Device(); + auto tp_group_name = nn::parallel::GetTensorParallelProcessGroupName(device.Rank().GlobalRank()); + auto tp_group = nn::parallel::ProcessGroupFactory::Instance()->Get(tp_group_name); + CHECK(tp_group != nullptr) << "TP group not found"; + + for (const auto &key : all_keys) { + + auto local_sd = Checkpoint::LoadStateDictFile(ckpt_dir / "model.ckpt"); + auto it = local_sd.find(key); + if (it == local_sd.end()) { + continue; + } + + auto local_tensor = it->second; + auto local_dims = local_tensor->Dims(); + + + std::vector gathered_dims = local_dims; + gathered_dims[0] *= tp_world_size; + + auto gathered_tensor + = std::make_shared(gathered_dims, local_tensor->Dtype(), local_tensor->GetDevice()); + + tp_group->AllGather(gathered_tensor, local_tensor, false); + merged_sd[key] = gathered_tensor; + } + + return merged_sd; +} + +// ----------------------------------------------------------------------------- + + +// ----------------------------------------------------------------------------- +static void LoadWithResharding(const std::filesystem::path &ckpt_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, + const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank) { + int saved_tp = metadata.parallel_config.tp_size; + int current_tp = nn::parallel::global::GetTensorParallelSize(); + bool tp_changed = saved_tp != current_tp; + + + auto load_result = checkpoint::LoadPlanner::PlanAndLoad(ckpt_dir, metadata); + auto &read_items = load_result.read_items; + auto &loaded_sd = load_result.loaded_sd; + + auto local_sd_it = loaded_sd.find("model.ckpt"); + CHECK(local_sd_it != loaded_sd.end()) << "model.ckpt not loaded"; + auto local_sd = std::move(local_sd_it->second); + + LOG(INFO) << "[CKPT] Loaded local shard: " << local_sd.size() << " tensors"; + + std::unordered_map> target_sd; + + if (tp_changed) { + + LOG(INFO) << "[CKPT] TP changed: " << saved_tp << " -> " << current_tp << ". Resharding..."; + + + auto full_sd = checkpoint::AllGatherFullModel(local_sd, saved_tp); + LOG(INFO) << "[CKPT] Gathered full model: " << full_sd.size() << " tensors"; + + auto reshard_plan = checkpoint::ReshardPlanner::ComputePlan( + read_items, metadata.parallel_config, + Checkpoint::CheckpointMetadata::ParallelConfig{current_tp, metadata.parallel_config.pp_size, + metadata.parallel_config.dp_size, + metadata.parallel_config.sp_size}, + tp_rank, pp_rank); + + target_sd = checkpoint::ReshardExecutor::Execute(reshard_plan, full_sd, tp_rank, pp_rank); + + LOG(INFO) << "[CKPT] Resharded to TP=" << current_tp << ", rank " << tp_rank << " has " << target_sd.size() + << " tensors"; + } else { + + target_sd = std::move(local_sd); + } + + + model.LoadStateDict(target_sd); + + + if (load_optimizer_state && optimizer != nullptr) { + auto opt_it = loaded_sd.find("optimizer.ckpt"); + if (opt_it != loaded_sd.end()) { + if (tp_changed) { + + auto opt_t = opt_it->second.find("adam.t"); + if (opt_t != opt_it->second.end()) { + optimizer->LoadStateDict({{"adam.t", opt_t->second}}); + LOG(WARNING) << "[CKPT] TP changed, optimizer m/v states reinitialized " + << "(only step counter loaded)"; + } + } else { + optimizer->LoadStateDict(opt_it->second); + } + } + } + + + state = Checkpoint::LoadTrainerStateFile(ckpt_dir / "trainer_state.json"); + + + state.tp_size = metadata.parallel_config.tp_size; + state.pp_size = metadata.parallel_config.pp_size; + state.sp_size = metadata.parallel_config.sp_size; + state.ddp_size = metadata.parallel_config.dp_size; + + + if (lr_scheduler != nullptr) { + auto lr_path = ckpt_dir / "lr_scheduler.ckpt"; + if (std::filesystem::exists(lr_path)) { + lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(lr_path)); + } else { + LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found at: " << lr_path; + } + } + + LOG(ERROR) << "[CKPT] LoadWithResharding done: global_step=" << state.global_step + << ", consumed_batches=" << state.consumed_batches; +} + +// ----------------------------------------------------------------------------- +// ResumeFromCheckpoint +// ----------------------------------------------------------------------------- ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { ResumeFromCheckpointResult result; if (args.resume_root.empty()) { @@ -31,38 +193,67 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & int sp_world_size = nn::parallel::global::GetSequenceParallelEnabled() ? tp_world_size : 1; int pp_world_size = nn::parallel::global::GetPipelineParallelSize(); - 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; + + std::filesystem::path ckpt_dir = args.resume_root; + + + auto metadata = Checkpoint::LoadMetadata(ckpt_dir); + if (!metadata.has_metadata) { + + if (args.rank.IsParallel()) { + const auto rank_dir = ckpt_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); + if (std::filesystem::exists(rank_dir)) { + ckpt_dir = rank_dir; + } } + + + Checkpoint::Load(ckpt_dir, *args.model, args.optimizer.get(), args.state, args.load_optimizer_state, + args.lr_scheduler.get()); + + result.global_step = static_cast(args.state.global_step); + + + 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.ddp_size, ddp_world_size) + << "DDP size mismatch: checkpoint has DDP=" << args.state.ddp_size << ", current=" << ddp_world_size; + CHECK_EQ(args.state.tp_size, tp_world_size) + << "TP size mismatch: checkpoint has TP=" << args.state.tp_size << ", current=" << tp_world_size; + + result.consumed_batches = static_cast(std::max(args.state.consumed_batches, 0)); + return result; } - Checkpoint::Load(resume_dir, *args.model, args.optimizer.get(), args.state, args.load_optimizer_state, - args.lr_scheduler.get()); + + LOG(INFO) << "[CKPT] Found metadata.json, iteration=" << metadata.iteration; + LOG(INFO) << "[CKPT] Saved parallel config: TP=" << metadata.parallel_config.tp_size + << ", PP=" << metadata.parallel_config.pp_size << ", DP=" << metadata.parallel_config.dp_size; + + + auto iter_dir = ckpt_dir / std::format("iter_{:07d}", static_cast(metadata.iteration)); + if (std::filesystem::exists(iter_dir / "model.ckpt")) { + ckpt_dir = iter_dir; + } + + + int global_rank = nn::parallel::global::GetGlobalProcRank(); + int dp_rank_dummy, tp_rank, pp_rank; + nn::parallel::global::GetCoordOf(global_rank, dp_rank_dummy, tp_rank, pp_rank); + int sp_rank = sp_world_size > 1 ? tp_rank : 0; + + LoadWithResharding(ckpt_dir, *args.model, args.optimizer.get(), args.state, args.load_optimizer_state, + args.lr_scheduler.get(), metadata, tp_rank, pp_rank, sp_rank); result.global_step = static_cast(args.state.global_step); + 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; result.consumed_micro_batches = static_cast(std::max(args.state.consumed_micro_batches, 0)); if (args.rank.IsMainRank()) { @@ -73,6 +264,9 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & return result; } +// ----------------------------------------------------------------------------- +// SaveCheckpoint +// ----------------------------------------------------------------------------- void SaveCheckpoint(const SaveCheckpointArgs &args) { const auto ckpt_start = std::chrono::high_resolution_clock::now(); @@ -89,26 +283,53 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { state.sp_size = args.sp_size; state.pp_size = args.pp_size; - Checkpoint::Save(args.save_dir, args.model, &args.optimizer, state, args.save_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(); + auto iter_dir = args.checkpoint_root_dir.empty() + ? args.save_dir + : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); + std::filesystem::create_directories(iter_dir); + + + if (args.rank.IsParallel() && !args.rank.IsMainRank()) { + auto rank_dir = iter_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); + std::filesystem::create_directories(rank_dir); - if (!args.rank.IsMainRank()) { + + Checkpoint::SaveTrainerStateFile(rank_dir / "trainer_state.json", state); + LOG(INFO) << "[CKPT] Non-main rank " << args.rank.GlobalRank() << " skipped model save"; return; } - LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", args.save_dir.string(), ckpt_ms); - // 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. + auto sharded_sd = args.model.ShardedStateDict("transformer"); + + + checkpoint::SavePlanner planner; + auto write_items = planner.Plan(sharded_sd, args.rank.GlobalRank()); + + + auto model_state_dict = args.model.StateDict(); + Checkpoint::SaveSharded(iter_dir, sharded_sd, write_items, model_state_dict, &args.optimizer, state, + args.save_optimizer_state, args.lr_scheduler, args.rank.GlobalRank()); + + + if (!args.checkpoint_root_dir.empty()) { + auto latest_file = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt"; + std::ofstream ofs(latest_file); + CHECK(ofs.is_open()) << "Failed to write latest checkpoint: " << latest_file; + ofs << args.global_step; + } + + const auto ckpt_end = std::chrono::high_resolution_clock::now(); + const double ckpt_ms = std::chrono::duration(ckpt_end - ckpt_start).count(); + + LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", iter_dir.string(), ckpt_ms); + + if (args.max_checkpoint_keep > 0 && std::filesystem::exists(args.checkpoint_root_dir)) { std::vector ckpts; for (const auto &entry : std::filesystem::directory_iterator(args.checkpoint_root_dir)) { - if (entry.is_directory() && entry.path().filename().string().starts_with("checkpoint_step_")) { + if (entry.is_directory() && entry.path().filename().string().starts_with("iter_")) { ckpts.push_back(entry.path()); } } diff --git a/infini_train/src/checkpoint/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc new file mode 100644 index 00000000..a4d1051a --- /dev/null +++ b/infini_train/src/checkpoint/load_planner.cc @@ -0,0 +1,146 @@ +#include "infini_train/include/checkpoint/load_planner.h" + +#include +#include +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/nn/parallel/global.h" + +namespace infini_train::checkpoint { + +// ----------------------------------------------------------------------------- +// ReadItem helpers +// ----------------------------------------------------------------------------- + +static std::string DataTypeToString(DataType dt) { + auto it = kDataTypeToDesc.find(dt); + if (it != kDataTypeToDesc.end()) { + return it->second; + } + return "fp32"; +} + +static DataType StringToDataType(const std::string &s) { + for (const auto &[dt, desc] : kDataTypeToDesc) { + if (desc == s) { + return dt; + } + } + return DataType::kFLOAT32; +} + +bool ReadItem::IsLocalRead(int /*current_tp_rank*/, int /*current_tp_size*/, int /*current_pp_rank*/, + int /*current_pp_size*/) const { + + + return true; +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- +std::vector LoadPlanner::Plan(const Checkpoint::CheckpointMetadata &metadata) { + std::vector items; + + for (const auto &entry : metadata.tensors) { + ReadItem item; + item.key = entry.key; + item.filename = entry.file; + item.dtype = StringToDataType(entry.dtype_str); + item.global_shape = entry.global_shape; + item.replica_id = 0; + item.byte_size = entry.byte_size; + + + + if (entry.key.find(".c_attn.") != std::string::npos || entry.key.find(".c_fc.") != std::string::npos + || entry.key.find(".c_fc2.") != std::string::npos || entry.key.find(".lm_head.") != std::string::npos + || entry.key.find(".wte.") != std::string::npos || entry.key.find(".output_layer.") != std::string::npos) { + ShardSpec spec; + spec.dim = 0; + spec.shard_count = metadata.parallel_config.tp_size; + spec.shard_index = 0; + spec.parallel_type = "tp"; + item.shard_specs.push_back(spec); + } else if (entry.key.find(".c_proj.") != std::string::npos) { + ShardSpec spec; + spec.dim = 1; + spec.shard_count = metadata.parallel_config.tp_size; + spec.shard_index = 0; + spec.parallel_type = "tp"; + item.shard_specs.push_back(spec); + } + + + items.push_back(std::move(item)); + } + + return items; +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- +std::unordered_map> +LoadPlanner::LoadFile(const std::filesystem::path &checkpoint_dir, const std::string &filename) { + auto path = checkpoint_dir / filename; + if (!std::filesystem::exists(path)) { + LOG(WARNING) << "[CKPT] File not found: " << path; + return {}; + } + return Checkpoint::LoadStateDictFile(path); +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- +LoadPlanner::PlanResult LoadPlanner::PlanAndLoad(const std::filesystem::path &checkpoint_dir, + const Checkpoint::CheckpointMetadata &metadata) { + + PlanResult result; + + + result.read_items = Plan(metadata); + + + std::unordered_set filenames; + for (const auto &item : result.read_items) { filenames.insert(item.filename); } + + for (const auto &filename : filenames) { + auto sd = LoadFile(checkpoint_dir, filename); + if (!sd.empty()) { + result.loaded_sd[filename] = std::move(sd); + LOG(INFO) << "[CKPT] Loaded " << filename << ": " << result.loaded_sd[filename].size() << " tensors"; + } + } + + + int missing = 0; + for (const auto &item : result.read_items) { + auto it = result.loaded_sd.find(item.filename); + if (it == result.loaded_sd.end()) { + LOG(WARNING) << "[CKPT] File " << item.filename << " not loaded, key=" << item.key; + missing++; + continue; + } + if (it->second.find(item.key) == it->second.end()) { + LOG(WARNING) << "[CKPT] Key " << item.key << " not found in " << item.filename; + missing++; + } + } + + if (missing > 0) { + LOG(WARNING) << "[CKPT] " << missing << " tensors missing from checkpoint files"; + } + + LOG(INFO) << "[CKPT] LoadPlanner: " << result.read_items.size() << " read items, " << result.loaded_sd.size() + << " files loaded"; + + 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 00000000..c670a41f --- /dev/null +++ b/infini_train/src/checkpoint/reshard.cc @@ -0,0 +1,379 @@ +#include "infini_train/include/checkpoint/reshard.h" + +#include +#include +#include +#include +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/save_planner.h" +#include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/nn/parallel/tensor_parallel.h" +#include "infini_train/include/nn/parallel/utils.h" +#include "infini_train/include/optimizer.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::checkpoint { + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- + + +static std::shared_ptr SliceTensorDim(const std::shared_ptr &full, int dim, int64_t offset, + int64_t size) { + auto dims = full->Dims(); + CHECK_LT(dim, static_cast(dims.size())) + << "dim " << dim << " out of range for tensor with " << dims.size() << " dims"; + CHECK_GE(offset, 0) << "negative offset: " << offset; + CHECK_LE(offset + size, dims[dim]) << "slice [" << offset << ":" << offset + size << ") exceeds dimension " + << dims[dim]; + + auto sliced = full->Slice(dim, offset, offset + size); + return std::make_shared(sliced->To(full->GetDevice())); +} + + +static std::shared_ptr SliceForTPKey(const std::shared_ptr &full_tensor, const std::string &key, + int tp_rank, int tp_world_size) { + + auto dims = full_tensor->Dims(); + auto part_type = ReshardPlanner::InferPartitionType(key); + + if (part_type == ReshardPlanner::PartitionType::kColumnParallel) { + CHECK_GE(dims.size(), 1u) << "Expected at least 1D tensor for " << key; + auto [start, size] = GetRankSliceRange(dims[0], tp_world_size, tp_rank); + return SliceTensorDim(full_tensor, 0, start, size); + } + + if (part_type == ReshardPlanner::PartitionType::kRowParallel) { + CHECK_GE(dims.size(), 2u) << "Expected 2D tensor for " << key; + auto [start, size] = GetRankSliceRange(dims[1], tp_world_size, tp_rank); + return SliceTensorDim(full_tensor, 1, start, size); + } + + + return std::make_shared(*full_tensor); +} + +// ----------------------------------------------------------------------------- +// ReshardPlanner +// ----------------------------------------------------------------------------- + +ReshardPlanner::PartitionType ReshardPlanner::InferPartitionType(const std::string &key) { + + + if (key.find(".c_attn.") != std::string::npos || key.find(".c_fc.") != std::string::npos + || key.find(".c_fc2.") != std::string::npos || key.find(".lm_head.") != std::string::npos + || key.find(".wte.") != std::string::npos || key.find(".output_layer.") != std::string::npos + || key.find("linear_qkv.") != std::string::npos || key.find("linear_fc1_gate.") != std::string::npos + || key.find("linear_fc1_up.") != std::string::npos + || key.find("linear_fc1.weight") != std::string::npos) { // fused gate+up + return PartitionType::kColumnParallel; + } + + + + if (key.find(".c_proj.") != std::string::npos || key.find("linear_proj.") != std::string::npos + || key.find("linear_fc2.") != std::string::npos) { + return PartitionType::kRowParallel; + } + + + return PartitionType::kReplicated; +} + +ReshardPlan ReshardPlanner::ComputePlan(const std::vector &read_items, + const Checkpoint::CheckpointMetadata::ParallelConfig &saved_config, + const Checkpoint::CheckpointMetadata::ParallelConfig ¤t_config, + int tp_rank, int pp_rank) { + + ReshardPlan plan; + plan.saved_tp_size = saved_config.tp_size; + plan.saved_pp_size = saved_config.pp_size; + plan.current_tp_size = current_config.tp_size; + plan.current_pp_size = current_config.pp_size; + plan.tp_changed = saved_config.tp_size != current_config.tp_size; + plan.pp_changed = saved_config.pp_size != current_config.pp_size; + + for (const auto &item : read_items) { + ReshardTensorPlan tp; + tp.key = item.key; + + auto part_type = InferPartitionType(item.key); + tp.is_replicated = (part_type == PartitionType::kReplicated); + + if (item.global_shape.empty()) { + LOG(WARNING) << "[CKPT] No global_shape for key: " << item.key; + tp.needs_reshard = false; + tp.target_shape = {}; + plan.tensors.push_back(std::move(tp)); + continue; + } + + if (tp.is_replicated || !plan.tp_changed) { + + tp.target_shape = item.global_shape; + tp.saved_local_shape = item.global_shape; + tp.needs_reshard = false; + plan.tensors.push_back(std::move(tp)); + continue; + } + + + tp.needs_reshard = true; + + int shard_dim = (part_type == PartitionType::kColumnParallel) ? 0 : 1; + int current_world = current_config.tp_size; + int64_t global_size = item.global_shape[shard_dim]; + + + auto [start, size] = GetRankSliceRange(global_size, current_world, tp_rank); + + + tp.target_shape = item.global_shape; + tp.target_shape[shard_dim] = size; + tp.saved_local_shape = item.global_shape; + + + ReshardSlice slice; + slice.dim = shard_dim; + slice.global_offset = start; + slice.size = size; + tp.slices.push_back(slice); + + plan.tensors.push_back(std::move(tp)); + } + + return plan; +} + +// ----------------------------------------------------------------------------- +// ReshardExecutor +// ----------------------------------------------------------------------------- + +std::unordered_map> +ReshardExecutor::Execute(const ReshardPlan &plan, + const std::unordered_map> &full_sd, int tp_rank, + int pp_rank) { + + (void)pp_rank; + + std::unordered_map> result; + int tp_world_size = nn::parallel::global::GetTensorParallelSize(); + + for (const auto &tp : plan.tensors) { + auto it = full_sd.find(tp.key); + if (it == full_sd.end()) { + LOG(WARNING) << "[CKPT] Reshard: key not found in full state_dict: " << tp.key; + continue; + } + + if (!tp.needs_reshard) { + + result[tp.key] = it->second; + continue; + } + + + auto tensor = it->second; + + if (!tp.slices.empty()) { + if (tp.slices.size() == 1) { + + const auto &slice = tp.slices[0]; + result[tp.key] = SliceTensorDim(tensor, slice.dim, slice.global_offset, slice.size); + } else { + + auto sliced = tensor; + for (const auto &slice : tp.slices) { + sliced = SliceTensorDim(sliced, slice.dim, slice.global_offset, slice.size); + } + result[tp.key] = sliced; + } + } else { + + result[tp.key] = SliceForTPKey(tensor, tp.key, tp_rank, tp_world_size); + } + } + + + for (const auto &[key, tensor] : full_sd) { + if (result.find(key) == result.end()) { + result[key] = tensor; + } + } + + LOG(INFO) << "[CKPT] ReshardExecutor: " << result.size() << " tensors ready for rank"; + + return result; +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- + +static const nn::parallel::ProcessGroup *GetTPGroup() { + auto device = Device(); + auto tp_group_name = nn::parallel::GetTensorParallelProcessGroupName(device.Rank().GlobalRank()); + return nn::parallel::ProcessGroupFactory::Instance()->Get(tp_group_name); +} + + + + + +std::unordered_map> +AllGatherFullModel(const std::unordered_map> &local_sd, int saved_tp_size) { + + auto tp_group = GetTPGroup(); + CHECK(tp_group != nullptr) << "TP process group not found, cannot gather for resharding"; + + std::unordered_map> full_sd; + + for (const auto &[key, local_tensor] : local_sd) { + auto local_dims = local_tensor->Dims(); + auto part_type = ReshardPlanner::InferPartitionType(key); + + + int shard_dim = (part_type == ReshardPlanner::PartitionType::kRowParallel) ? 1 : 0; + + if (shard_dim == 0) { + + std::vector full_dims = local_dims; + full_dims[0] *= saved_tp_size; + + auto full_tensor = std::make_shared(full_dims, local_tensor->Dtype(), local_tensor->GetDevice()); + tp_group->AllGather(full_tensor, local_tensor, false); + full_sd[key] = full_tensor; + + } else { + + + + auto local_transposed = local_tensor->Transpose(0, 1); + + + std::vector gathered_dims = local_transposed->Dims(); + gathered_dims[0] *= saved_tp_size; + + auto gathered = std::make_shared(gathered_dims, local_tensor->Dtype(), local_tensor->GetDevice()); + tp_group->AllGather(gathered, local_transposed, false); + + + + auto full_tensor = std::make_shared(gathered->Transpose(0, 1)->To(local_tensor->GetDevice())); + full_sd[key] = full_tensor; + } + } + + return full_sd; +} + +// ----------------------------------------------------------------------------- + +// ----------------------------------------------------------------------------- + +void ReshardAndLoad(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, + const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank) { + + (void)sp_rank; + + int saved_tp = metadata.parallel_config.tp_size; + int current_tp = nn::parallel::global::GetTensorParallelSize(); + bool tp_changed = saved_tp != current_tp; + + + auto load_result = LoadPlanner::PlanAndLoad(checkpoint_dir, metadata); + auto &read_items = load_result.read_items; + auto &loaded_sd = load_result.loaded_sd; + + + auto local_sd_it = loaded_sd.find("model.ckpt"); + CHECK(local_sd_it != loaded_sd.end()) << "model.ckpt not loaded"; + auto local_sd = std::move(local_sd_it->second); + + LOG(INFO) << "[CKPT] Loaded local shard: " << local_sd.size() << " tensors"; + + std::unordered_map> target_sd; + + if (tp_changed) { + + LOG(INFO) << "[CKPT] TP changed: " << saved_tp << " -> " << current_tp + << ". Gathering full model for resharding..."; + + auto full_sd = AllGatherFullModel(local_sd, saved_tp); + LOG(INFO) << "[CKPT] Gathered full model: " << full_sd.size() << " tensors"; + + + auto reshard_plan + = ReshardPlanner::ComputePlan(read_items, metadata.parallel_config, + Checkpoint::CheckpointMetadata::ParallelConfig{ + current_tp, metadata.parallel_config.pp_size, + metadata.parallel_config.dp_size, metadata.parallel_config.sp_size}, + tp_rank, pp_rank); + + + target_sd = ReshardExecutor::Execute(reshard_plan, full_sd, tp_rank, pp_rank); + + LOG(INFO) << "[CKPT] Resharded to TP=" << current_tp << ", rank " << tp_rank << " has " << target_sd.size() + << " tensors"; + } else { + + target_sd = std::move(local_sd); + } + + + model.LoadStateDict(target_sd); + + + if (load_optimizer_state && optimizer != nullptr) { + auto opt_it = loaded_sd.find("optimizer.ckpt"); + if (opt_it != loaded_sd.end()) { + if (tp_changed) { + + auto opt_t = opt_it->second.find("adam.t"); + if (opt_t != opt_it->second.end()) { + optimizer->LoadStateDict({{"adam.t", opt_t->second}}); + LOG(WARNING) << "[CKPT] TP changed, optimizer m/v states reinitialized " + << "(only step counter loaded)"; + } + } else { + optimizer->LoadStateDict(opt_it->second); + } + } + } + + + state = Checkpoint::LoadTrainerStateFile(checkpoint_dir / "trainer_state.json"); + + + state.tp_size = saved_tp; + state.pp_size = metadata.parallel_config.pp_size; + state.sp_size = metadata.parallel_config.sp_size; + state.ddp_size = metadata.parallel_config.dp_size; + + + if (lr_scheduler != nullptr) { + auto lr_path = checkpoint_dir / "lr_scheduler.ckpt"; + if (std::filesystem::exists(lr_path)) { + lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(lr_path)); + } else { + LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found: " << lr_path; + } + } + + LOG(ERROR) << "[CKPT] ReshardAndLoad done: global_step=" << state.global_step + << ", consumed_batches=" << state.consumed_batches << ", TP=" << current_tp + << ", tensors=" << target_sd.size(); +} + +} // 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 00000000..284d3d0c --- /dev/null +++ b/infini_train/src/checkpoint/save_planner.cc @@ -0,0 +1,32 @@ +#include "infini_train/include/checkpoint/save_planner.h" + +namespace infini_train::checkpoint { + +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.shard_specs = info.shard_specs; + item.replica_id = info.replica_id; + 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 fb9e53e1..f4ba274c 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -182,6 +182,45 @@ std::unordered_map> Module::StateDict() con return state; } +checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix, + const std::vector &parent_shards) const { + checkpoint::ShardedStateDict sd; + + for (auto &[name, param] : parameters_) { + checkpoint::ShardedTensorInfo info; + info.key = prefix.empty() ? name : prefix + "." + name; + info.dtype = param->Dtype(); + info.global_shape = param->Dims(); + info.local_shape = param->Dims(); + info.shard_specs = parent_shards; + info.replica_id = 0; + sd.tensors[info.key] = std::move(info); + } + + for (auto &[name, buffer] : buffers_) { + checkpoint::ShardedTensorInfo info; + info.key = prefix.empty() ? name : prefix + "." + name; + info.dtype = buffer->Dtype(); + info.global_shape = buffer->Dims(); + info.local_shape = buffer->Dims(); + info.shard_specs = parent_shards; + info.replica_id = 0; + 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, parent_shards); + 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/parallel/tensor_parallel.cc b/infini_train/src/nn/parallel/tensor_parallel.cc index a653191b..cf858340 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -284,6 +284,47 @@ 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 std::vector &parent_shards) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + // Weight is split along output dimension (dim=0) + auto tp_shards = parent_shards; + tp_shards.push_back({ + .dim = 0, + .shard_count = tp_size, + .shard_index = tp_rank, + .parallel_type = "tp", + }); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensorInfo w; + w.key = prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {output_size_per_partition_ * tp_size, weight->Dims()[1]}; + w.local_shape = weight->Dims(); + w.shard_specs = tp_shards; + w.replica_id = 0; + sd.tensors[w.key] = std::move(w); + + // Bias is also split along dim=0 + if (bias_) { + auto &bias = parameter(kParamBiasName); + checkpoint::ShardedTensorInfo b; + b.key = prefix + "." + kParamBiasName; + b.dtype = bias->Dtype(); + b.global_shape = {static_cast(output_size_per_partition_ * tp_size)}; + b.local_shape = bias->Dims(); + b.shard_specs = tp_shards; + b.replica_id = 0; + 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 +380,47 @@ 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 std::vector &parent_shards) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + // Weight is split along input dimension (dim=1) + auto tp_shards = parent_shards; + tp_shards.push_back({ + .dim = 1, + .shard_count = tp_size, + .shard_index = tp_rank, + .parallel_type = "tp", + }); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensorInfo w; + w.key = prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {weight->Dims()[0], input_size_per_partition_ * tp_size}; + w.local_shape = weight->Dims(); + w.shard_specs = tp_shards; + w.replica_id = 0; + sd.tensors[w.key] = std::move(w); + + // Bias is NOT sharded in RowParallelLinear + if (bias_) { + auto &bias = parameter(kParamBiasName); + checkpoint::ShardedTensorInfo b; + b.key = prefix + "." + kParamBiasName; + b.dtype = bias->Dtype(); + b.global_shape = bias->Dims(); + b.local_shape = bias->Dims(); + b.shard_specs = parent_shards; // no TP shard for bias + b.replica_id = 0; + 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), vocab_size_global_(num_embeddings), embedding_dim_(embedding_dim), @@ -396,6 +478,33 @@ VocabParallelEmbedding::Forward(const std::vector> &inpu return {output}; } +checkpoint::ShardedStateDict +VocabParallelEmbedding::ShardedStateDict(const std::string &prefix, + const std::vector &parent_shards) const { + checkpoint::ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto tp_shards = parent_shards; + tp_shards.push_back({ + .dim = 0, + .shard_count = tp_size, + .shard_index = tp_rank, + .parallel_type = "tp", + }); + + auto &weight = parameter(kParamWeightName); + checkpoint::ShardedTensorInfo w; + w.key = prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {vocab_size_global_, embedding_dim_}; + w.local_shape = weight->Dims(); + w.shard_specs = tp_shards; + w.replica_id = 0; + 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}"; @@ -466,7 +575,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) + auto softmax_local = exp_local->Div(sum_exp->Unsqueeze(-1)); // 5. Perform allreduce to get global predicted_logit @@ -480,7 +589,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) + // TODO(zbl): adjust smoothing coef according to vocab_size_original if (label_smoothing_ > 0.0f) { // mean_logp over *valid tokens only*: From 8fcb2a5423d7d7099864c96a92758f479cbe8f91 Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Tue, 21 Jul 2026 10:23:39 +0000 Subject: [PATCH 2/3] feat: complete distributed checkpoint resharding --- infini_train/include/checkpoint/checkpoint.h | 26 +- .../include/checkpoint/checkpoint_manager.h | 1 - .../include/checkpoint/load_planner.h | 47 +- .../include/checkpoint/load_strategy.h | 30 ++ infini_train/include/checkpoint/reshard.h | 84 +--- .../include/checkpoint/save_planner.h | 23 +- infini_train/include/checkpoint/shard_spec.h | 42 +- infini_train/include/nn/modules/module.h | 12 +- .../transformer/causal_self_attention.h | 2 + .../nn/modules/transformer/transformer.h | 5 + .../parallel/ddp/distributed_data_parallel.h | 2 + .../nn/parallel/pp/pipeline_parallel.h | 6 + .../include/nn/parallel/tensor_parallel.h | 12 +- infini_train/src/checkpoint/checkpoint.cc | 337 +++++++++++--- .../src/checkpoint/checkpoint_manager.cc | 411 ++++++------------ infini_train/src/checkpoint/load_planner.cc | 307 ++++++++----- infini_train/src/checkpoint/load_strategy.cc | 129 ++++++ infini_train/src/checkpoint/reshard.cc | 402 ++--------------- infini_train/src/checkpoint/save_planner.cc | 46 +- infini_train/src/nn/modules/module.cc | 40 +- .../transformer/causal_self_attention.cc | 34 ++ .../src/nn/modules/transformer/transformer.cc | 77 ++++ .../parallel/ddp/distributed_data_parallel.cc | 5 + .../src/nn/parallel/pp/pipeline_parallel.cc | 17 + .../src/nn/parallel/tensor_parallel.cc | 77 ++-- .../test_checkpoint_serialization.cc | 382 ++++++++++++++++ tests/checkpoint/test_optimizer_state.cc | 19 + .../test_transformer_architecture.cc | 4 + 28 files changed, 1556 insertions(+), 1023 deletions(-) create mode 100644 infini_train/include/checkpoint/load_strategy.h create mode 100644 infini_train/src/checkpoint/load_strategy.cc diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index 8d493479..51959dc0 100644 --- a/infini_train/include/checkpoint/checkpoint.h +++ b/infini_train/include/checkpoint/checkpoint.h @@ -45,8 +45,8 @@ class Checkpoint { 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 Optimizer *optimizer, const TrainerState &state, bool save_optimizer_state, - const LRScheduler *lr_scheduler, int global_rank); + 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); @@ -69,10 +69,16 @@ class Checkpoint { 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; + int replica_id = 0; std::vector stored_on_ranks; + int pp_rank = 0; }; std::vector tensors; @@ -80,18 +86,26 @@ 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: - 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 3747790f..150bd6a8 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 index bcaa5982..7011cc73 100644 --- a/infini_train/include/checkpoint/load_planner.h +++ b/infini_train/include/checkpoint/load_planner.h @@ -1,49 +1,52 @@ #pragma once #include -#include -#include +#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" -#include "infini_train/include/tensor.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; - std::vector shard_specs; - int replica_id = 0; 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; +}; - bool IsLocalRead(int current_tp_rank, int current_tp_size, int current_pp_rank, int current_pp_size) const; +// 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: - struct PlanResult { - std::vector read_items; - std::unordered_map>> loaded_sd; - }; - - - static PlanResult PlanAndLoad(const std::filesystem::path &checkpoint_dir, - const Checkpoint::CheckpointMetadata &metadata); - - - static std::vector Plan(const Checkpoint::CheckpointMetadata &metadata); - - - static std::unordered_map> - LoadFile(const std::filesystem::path &checkpoint_dir, const std::string &filename); + // 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 00000000..f5f6e2c0 --- /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 index c457b4ad..70970538 100644 --- a/infini_train/include/checkpoint/reshard.h +++ b/infini_train/include/checkpoint/reshard.h @@ -1,20 +1,12 @@ #pragma once -#include #include -#include -#include -#include -#include #include "infini_train/include/checkpoint/checkpoint.h" -#include "infini_train/include/checkpoint/load_planner.h" -#include "infini_train/include/checkpoint/shard_spec.h" -#include "infini_train/include/tensor.h" namespace infini_train { -class Optimizer; class LRScheduler; +class Optimizer; namespace nn { class Module; } @@ -22,75 +14,9 @@ class Module; namespace infini_train::checkpoint { - -struct ReshardSlice { - int dim; - int64_t global_offset; - int64_t size; -}; - - -struct ReshardTensorPlan { - std::string key; - std::vector saved_local_shape; - std::vector target_shape; - std::vector slices; - bool needs_reshard = false; - bool is_replicated = false; -}; - - -struct ReshardPlan { - std::vector tensors; - int saved_tp_size = 1; - int saved_pp_size = 1; - int current_tp_size = 1; - int current_pp_size = 1; - bool tp_changed = false; - bool pp_changed = false; -}; - - - -class ReshardPlanner { -public: - - static ReshardPlan ComputePlan(const std::vector &read_items, - const Checkpoint::CheckpointMetadata::ParallelConfig &saved_config, - const Checkpoint::CheckpointMetadata::ParallelConfig ¤t_config, int tp_rank, - int pp_rank); - - - enum class PartitionType { - kColumnParallel, - kRowParallel, - kReplicated, - }; - - static PartitionType InferPartitionType(const std::string &key); -}; - - -class ReshardExecutor { -public: - - - - - static std::unordered_map> - Execute(const ReshardPlan &plan, const std::unordered_map> &full_sd, - int tp_rank, int pp_rank); -}; - - - - -std::unordered_map> -AllGatherFullModel(const std::unordered_map> &local_sd, int saved_tp_size); - - -void ReshardAndLoad(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, - TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, - const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank); +// 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, bool load_optimizer_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 192ec0f1..e7af9ae1 100644 --- a/infini_train/include/checkpoint/save_planner.h +++ b/infini_train/include/checkpoint/save_planner.h @@ -1,34 +1,45 @@ #pragma once #include +#include #include +#include #include #include "infini_train/include/checkpoint/shard_spec.h" #include "infini_train/include/datatype.h" -namespace infini_train::checkpoint { +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; - uint64_t byte_size = 0; + 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 shard_specs; + std::vector global_offset; + std::vector axis_fragmentations; int replica_id = 0; 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); } @@ -57,7 +68,7 @@ inline uint64_t TensorByteSize(DataType dtype, const std::vector &shape } } - +// 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; diff --git a/infini_train/include/checkpoint/shard_spec.h b/infini_train/include/checkpoint/shard_spec.h index 0cb16d11..d6c48951 100644 --- a/infini_train/include/checkpoint/shard_spec.h +++ b/infini_train/include/checkpoint/shard_spec.h @@ -9,39 +9,39 @@ namespace infini_train::checkpoint { +struct ShardSegment { + int64_t global_offset = 0; + int64_t local_offset = 0; + int64_t length = 0; -struct ShardSpec { - int dim = -1; - int shard_count = 1; - int shard_index = 0; - std::string parallel_type; // "tp" / "pp" / "dp" / "sp" - - bool operator==(const ShardSpec &other) const { - return dim == other.dim && shard_count == other.shard_count && shard_index == other.shard_index - && parallel_type == other.parallel_type; - } - bool operator!=(const ShardSpec &other) const { return !(*this == other); } + bool operator==(const ShardSegment &other) const = default; }; - -struct ShardedTensorInfo { - std::string key; // "transformer.h.0.attn.c_attn.weight" +// 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 shard_specs; + 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; int replica_id = 0; - bool operator==(const ShardedTensorInfo &other) const { - return key == other.key && dtype == other.dtype && global_shape == other.global_shape - && local_shape == other.local_shape && shard_specs == other.shard_specs; + 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 && replica_id == other.replica_id; } }; - struct ShardedStateDict { - std::map tensors; - + std::map tensors; void Merge(ShardedStateDict &&other) { for (auto &[key, info] : other.tensors) { tensors.emplace(std::move(key), std::move(info)); } diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 070dcb8d..880f26da 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -50,7 +50,7 @@ class Module : public std::enable_shared_from_this { // TODO: Change return type to filterable iterator (like PyTorch's named_parameters with prefix matching) virtual std::vector> Parameters() const; - 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); @@ -62,16 +62,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; - - virtual checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", - const std::vector &parent_shards - = {}) 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 60373414..2769dd1d 100644 --- a/infini_train/include/nn/modules/transformer/causal_self_attention.h +++ b/infini_train/include/nn/modules/transformer/causal_self_attention.h @@ -21,6 +21,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 0471c32f..455f37c2 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 823ae82b..1ad2d44c 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -30,6 +30,8 @@ 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; DistributedDataParallelConfig ddp_config() const { return ddp_config_; } diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 25939bdc..58f48cd5 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 c7a17429..6847b5fb 100644 --- a/infini_train/include/nn/parallel/tensor_parallel.h +++ b/infini_train/include/nn/parallel/tensor_parallel.h @@ -38,9 +38,7 @@ class ColumnParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", - const std::vector &parent_shards - = {}) const override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; protected: bool bias_ = true; @@ -71,9 +69,7 @@ class RowParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", - const std::vector &parent_shards - = {}) const override; + checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "") const override; protected: bool bias_ = true; @@ -94,9 +90,7 @@ class VocabParallelEmbedding : public nn::CloneableModule> Forward(const std::vector> &input_tensors) override; - checkpoint::ShardedStateDict ShardedStateDict(const std::string &prefix = "", - const std::vector &parent_shards - = {}) const 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 c997c948..3abe8b2c 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 @@ -14,6 +15,7 @@ #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" @@ -244,13 +246,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)); @@ -270,8 +274,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) { @@ -358,6 +368,14 @@ void Checkpoint::SaveTrainerStateFile(const std::filesystem::path &path, const T 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); @@ -368,16 +386,8 @@ Checkpoint::LoadStateDictFile(const std::filesystem::path &path) { return LoadStateDict(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); -} - // ----------------------------------------------------------------------------- - +// Save local shards and a temporary rank manifest from a ShardedStateDict. // ----------------------------------------------------------------------------- static std::string DataTypeToString(DataType dt) { @@ -392,57 +402,50 @@ 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 Optimizer *optimizer, const TrainerState &state, bool save_optimizer_state, - const LRScheduler *lr_scheduler, int global_rank) { + 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 replica-zero model tensors. { std::unordered_map> filtered_sd; for (const auto &[key, info] : sharded_sd.tensors) { if (info.replica_id != 0) { continue; } - + // Optimizer tensors are serialized separately. if (key.starts_with("adam.")) { continue; } - - auto it = state_dict.find(key); + // 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()) { - SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); + model_file_index = SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); } } - - if (save_optimizer_state && optimizer != nullptr) { - auto opt_state = optimizer->StateDict(); - if (!opt_state.empty()) { - SaveStateDict(checkpoint_dir / "optimizer.ckpt", opt_state); - } + // Save the rank-local optimizer state. + if (!optimizer_state.empty()) { + optimizer_file_index = SaveStateDict(checkpoint_dir / "optimizer.ckpt", optimizer_state); } - - if (lr_scheduler != nullptr) { - SaveLRSchedulerState(checkpoint_dir / "lr_scheduler.ckpt", lr_scheduler->StateDict()); - } - - - SaveTrainerState(checkpoint_dir / "trainer_state.json", 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\": 1,\n"; + ofs << " \"version\": 3,\n"; ofs << " \"format\": \"infinitrain_sharded\",\n"; ofs << " \"iteration\": " << state.global_step << ",\n"; ofs << " \"parallel_config\": {\n"; @@ -460,16 +463,22 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, ofs << " },\n"; ofs << " \"tensors\": [\n"; - for (size_t i = 0; i < write_items.size(); ++i) { - auto &item = write_items[i]; - if (item.replica_id != 0) { - continue; - } - - auto it = sharded_sd.tensors.find(item.key); - if (it == sharded_sd.tensors.end()) { - continue; + std::vector emitted_items; + for (const auto &item : write_items) { + if (item.replica_id == 0 && 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"; @@ -481,28 +490,40 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, for (size_t j = 0; j < gs.size(); ++j) { ofs << gs[j] << (j + 1 < gs.size() ? ", " : ""); } ofs << "],\n"; - // shard_specs - if (!it->second.shard_specs.empty()) { - ofs << " \"shard_specs\": [\n"; - for (size_t s = 0; s < it->second.shard_specs.size(); ++s) { - auto &spec = it->second.shard_specs[s]; - ofs << " {\"dim\": " << spec.dim << ", \"shard_count\": " << spec.shard_count - << ", \"shard_index\": " << spec.shard_index << ", \"parallel_type\": \"" << spec.parallel_type - << "\"}"; - if (s + 1 < it->second.shard_specs.size()) { - ofs << ","; - } - ofs << "\n"; - } - 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\": " << item.offset << ",\n"; + ofs << " \"offset\": " << storage.data_offset << ",\n"; ofs << " \"byte_size\": " << item.byte_size << ",\n"; + ofs << " \"replica_id\": " << item.replica_id << ",\n"; + ofs << " \"pp_rank\": " << pp_rank << ",\n"; ofs << " \"stored_on_ranks\": [" << global_rank << "]\n"; ofs << " }"; - if (i + 1 < write_items.size()) { + if (i + 1 < emitted_items.size()) { ofs << ","; } ofs << "\n"; @@ -517,10 +538,7 @@ void Checkpoint::SaveSharded(const std::filesystem::path &checkpoint_dir, 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); @@ -538,8 +556,8 @@ static std::string ExtractJsonString(const std::string &obj, const std::string & return obj.substr(q1 + 1, q2 - q1 - 1); } -Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::path &checkpoint_dir) { - CheckpointMetadata meta; +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; @@ -558,7 +576,7 @@ Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::p 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; @@ -581,10 +599,23 @@ Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::p } 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) { - size_t obj_end = tensor_block.find('}', obj_pos); + 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; } @@ -597,6 +628,8 @@ Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::p entry.dtype_str = ExtractJsonString(obj, "dtype"); entry.offset = ExtractNumberField(obj, "offset", 0); entry.byte_size = ExtractNumberField(obj, "byte_size", 0); + entry.replica_id = ExtractNumberField(obj, "replica_id", 0); + entry.pp_rank = ExtractNumberField(obj, "pp_rank", 0); // global_shape: [x, y, z] auto gs_pos = obj.find("\"global_shape\""); @@ -615,6 +648,89 @@ Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::p } } + 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; } @@ -622,4 +738,91 @@ Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::p 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 << " \"replica_id\": " << tensor.replica_id << ",\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 fd843b18..5b2faddc 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -2,185 +2,73 @@ #include #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/load_planner.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/process_group.h" -#include "infini_train/include/nn/parallel/rank.h" -#include "infini_train/include/nn/parallel/tensor_parallel.h" -#include "infini_train/include/nn/parallel/utils.h" -#include "infini_train/include/tensor.h" using namespace infini_train; namespace nn = infini_train::nn; -// ----------------------------------------------------------------------------- +namespace { -// ----------------------------------------------------------------------------- -static std::unordered_map> LoadRankShards(const std::filesystem::path &ckpt_dir, - int tp_world_size) { - std::unordered_map> merged_sd; - std::unordered_set all_keys; - - - for (int r = 0; r < tp_world_size; ++r) { - auto rank_dir = ckpt_dir / std::format("rank_{:06d}", r); - auto model_path = rank_dir / "model.ckpt"; - if (!std::filesystem::exists(model_path)) { - - model_path = ckpt_dir / "model.ckpt"; - } - if (!std::filesystem::exists(model_path)) { - LOG(WARNING) << "[CKPT] No model.ckpt found for rank " << r; - continue; - } - auto sd = Checkpoint::LoadStateDictFile(model_path); - for (const auto &[key, tensor] : sd) { all_keys.insert(key); } - } - - - auto device = Device(); - auto tp_group_name = nn::parallel::GetTensorParallelProcessGroupName(device.Rank().GlobalRank()); - auto tp_group = nn::parallel::ProcessGroupFactory::Instance()->Get(tp_group_name); - CHECK(tp_group != nullptr) << "TP group not found"; - - for (const auto &key : all_keys) { - - auto local_sd = Checkpoint::LoadStateDictFile(ckpt_dir / "model.ckpt"); - auto it = local_sd.find(key); - if (it == local_sd.end()) { - continue; - } - - auto local_tensor = it->second; - auto local_dims = local_tensor->Dims(); - - - std::vector gathered_dims = local_dims; - gathered_dims[0] *= tp_world_size; - - auto gathered_tensor - = std::make_shared(gathered_dims, local_tensor->Dtype(), local_tensor->GetDevice()); - - tp_group->AllGather(gathered_tensor, local_tensor, false); - merged_sd[key] = gathered_tensor; +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; } - - return merged_sd; + 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; } -// ----------------------------------------------------------------------------- - - -// ----------------------------------------------------------------------------- -static void LoadWithResharding(const std::filesystem::path &ckpt_dir, nn::Module &model, Optimizer *optimizer, - TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, - const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank) { - int saved_tp = metadata.parallel_config.tp_size; - int current_tp = nn::parallel::global::GetTensorParallelSize(); - bool tp_changed = saved_tp != current_tp; - - - auto load_result = checkpoint::LoadPlanner::PlanAndLoad(ckpt_dir, metadata); - auto &read_items = load_result.read_items; - auto &loaded_sd = load_result.loaded_sd; - - auto local_sd_it = loaded_sd.find("model.ckpt"); - CHECK(local_sd_it != loaded_sd.end()) << "model.ckpt not loaded"; - auto local_sd = std::move(local_sd_it->second); - - LOG(INFO) << "[CKPT] Loaded local shard: " << local_sd.size() << " tensors"; - - std::unordered_map> target_sd; - - if (tp_changed) { - - LOG(INFO) << "[CKPT] TP changed: " << saved_tp << " -> " << current_tp << ". Resharding..."; - - - auto full_sd = checkpoint::AllGatherFullModel(local_sd, saved_tp); - LOG(INFO) << "[CKPT] Gathered full model: " << full_sd.size() << " tensors"; - - auto reshard_plan = checkpoint::ReshardPlanner::ComputePlan( - read_items, metadata.parallel_config, - Checkpoint::CheckpointMetadata::ParallelConfig{current_tp, metadata.parallel_config.pp_size, - metadata.parallel_config.dp_size, - metadata.parallel_config.sp_size}, - tp_rank, pp_rank); - - target_sd = checkpoint::ReshardExecutor::Execute(reshard_plan, full_sd, tp_rank, pp_rank); - - LOG(INFO) << "[CKPT] Resharded to TP=" << current_tp << ", rank " << tp_rank << " has " << target_sd.size() - << " tensors"; - } else { - - target_sd = std::move(local_sd); - } - - - model.LoadStateDict(target_sd); - - - if (load_optimizer_state && optimizer != nullptr) { - auto opt_it = loaded_sd.find("optimizer.ckpt"); - if (opt_it != loaded_sd.end()) { - if (tp_changed) { - - auto opt_t = opt_it->second.find("adam.t"); - if (opt_t != opt_it->second.end()) { - optimizer->LoadStateDict({{"adam.t", opt_t->second}}); - LOG(WARNING) << "[CKPT] TP changed, optimizer m/v states reinitialized " - << "(only step counter loaded)"; +void WaitForWriterManifests(const std::filesystem::path &staging_root, int tp_size, int pp_size) { + 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); + if (!std::filesystem::exists(staging_root / std::format("rank_{:06d}", rank) / "metadata.json")) { + ready = false; + break; } - } else { - optimizer->LoadStateDict(opt_it->second); } } - } - - - state = Checkpoint::LoadTrainerStateFile(ckpt_dir / "trainer_state.json"); - - - state.tp_size = metadata.parallel_config.tp_size; - state.pp_size = metadata.parallel_config.pp_size; - state.sp_size = metadata.parallel_config.sp_size; - state.ddp_size = metadata.parallel_config.dp_size; - - - if (lr_scheduler != nullptr) { - auto lr_path = ckpt_dir / "lr_scheduler.ckpt"; - if (std::filesystem::exists(lr_path)) { - lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(lr_path)); - } else { - LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found at: " << lr_path; + 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)); } +} - LOG(ERROR) << "[CKPT] LoadWithResharding done: global_step=" << state.global_step - << ", consumed_batches=" << state.consumed_batches; +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)); + } } -// ----------------------------------------------------------------------------- -// ResumeFromCheckpoint -// ----------------------------------------------------------------------------- +} // namespace + ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { ResumeFromCheckpointResult result; if (args.resume_root.empty()) { @@ -188,155 +76,126 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & return result; } - 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(); - - - std::filesystem::path ckpt_dir = args.resume_root; + 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.load_optimizer_state, args.lr_scheduler.get(), metadata); - auto metadata = Checkpoint::LoadMetadata(ckpt_dir); - if (!metadata.has_metadata) { - - if (args.rank.IsParallel()) { - const auto rank_dir = ckpt_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); - if (std::filesystem::exists(rank_dir)) { - ckpt_dir = rank_dir; - } - } - - - Checkpoint::Load(ckpt_dir, *args.model, args.optimizer.get(), args.state, args.load_optimizer_state, - args.lr_scheduler.get()); - - result.global_step = static_cast(args.state.global_step); - - - 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.ddp_size, ddp_world_size) - << "DDP size mismatch: checkpoint has DDP=" << args.state.ddp_size << ", current=" << ddp_world_size; - CHECK_EQ(args.state.tp_size, tp_world_size) - << "TP size mismatch: checkpoint has TP=" << args.state.tp_size << ", current=" << tp_world_size; - - result.consumed_batches = static_cast(std::max(args.state.consumed_batches, 0)); - return result; - } - - - LOG(INFO) << "[CKPT] Found metadata.json, iteration=" << metadata.iteration; - LOG(INFO) << "[CKPT] Saved parallel config: TP=" << metadata.parallel_config.tp_size - << ", PP=" << metadata.parallel_config.pp_size << ", DP=" << metadata.parallel_config.dp_size; - - - auto iter_dir = ckpt_dir / std::format("iter_{:07d}", static_cast(metadata.iteration)); - if (std::filesystem::exists(iter_dir / "model.ckpt")) { - ckpt_dir = iter_dir; - } - - - int global_rank = nn::parallel::global::GetGlobalProcRank(); - int dp_rank_dummy, tp_rank, pp_rank; - nn::parallel::global::GetCoordOf(global_rank, dp_rank_dummy, tp_rank, pp_rank); - int sp_rank = sp_world_size > 1 ? tp_rank : 0; - - LoadWithResharding(ckpt_dir, *args.model, args.optimizer.get(), args.state, args.load_optimizer_state, - args.lr_scheduler.get(), metadata, tp_rank, pp_rank, sp_rank); - + 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_embd, args.model_config.n_embd); result.global_step = static_cast(args.state.global_step); - - - 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_embd, args.model_config.n_embd) - << "n_embd mismatch: ckpt=" << args.state.n_embd << ", config=" << args.model_config.n_embd; - result.consumed_micro_batches = static_cast(std::max(args.state.consumed_micro_batches, 0)); if (args.rank.IsMainRank()) { LOG(INFO) << std::format("Resume training from step {}, consumed_micro_batches {}", args.state.global_step, args.state.consumed_micro_batches); } - return result; } -// ----------------------------------------------------------------------------- -// SaveCheckpoint -// ----------------------------------------------------------------------------- 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_micro_batches = static_cast(args.consumed_micro_batches); - state.n_layer = args.n_layer; - state.n_head = args.n_head; - state.n_kv_head = args.n_kv_head; - 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; - - - auto iter_dir = args.checkpoint_root_dir.empty() - ? args.save_dir - : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); - std::filesystem::create_directories(iter_dir); - - - if (args.rank.IsParallel() && !args.rank.IsMainRank()) { - auto rank_dir = iter_dir / std::format("rank_{:06d}", args.rank.GlobalRank()); - std::filesystem::create_directories(rank_dir); - - - Checkpoint::SaveTrainerStateFile(rank_dir / "trainer_state.json", state); - LOG(INFO) << "[CKPT] Non-main rank " << args.rank.GlobalRank() << " skipped model save"; + const auto checkpoint_start = std::chrono::high_resolution_clock::now(); + TrainerState state{.global_step = args.global_step, + .consumed_micro_batches = static_cast(args.consumed_micro_batches), + .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); + + 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; } - - auto sharded_sd = args.model.ShardedStateDict("transformer"); - - - checkpoint::SavePlanner planner; - auto write_items = planner.Plan(sharded_sd, args.rank.GlobalRank()); - - - auto model_state_dict = args.model.StateDict(); - Checkpoint::SaveSharded(iter_dir, sharded_sd, write_items, model_state_dict, &args.optimizer, state, - args.save_optimizer_state, args.lr_scheduler, args.rank.GlobalRank()); - - - if (!args.checkpoint_root_dir.empty()) { - auto latest_file = args.checkpoint_root_dir / "latest_checkpointed_iteration.txt"; - std::ofstream ofs(latest_file); - CHECK(ofs.is_open()) << "Failed to write latest checkpoint: " << latest_file; - ofs << args.global_step; + 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.save_optimizer_state) { + 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_root = iteration_dir / ".metadata_tmp"; + 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); - const auto ckpt_end = std::chrono::high_resolution_clock::now(); - const double ckpt_ms = std::chrono::duration(ckpt_end - ckpt_start).count(); - - LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", iter_dir.string(), ckpt_ms); + 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); + 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.max_checkpoint_keep > 0 && std::filesystem::exists(args.checkpoint_root_dir)) { - std::vector ckpts; + 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("iter_")) { - ckpts.push_back(entry.path()); + 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); } diff --git a/infini_train/src/checkpoint/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc index a4d1051a..334268f4 100644 --- a/infini_train/src/checkpoint/load_planner.cc +++ b/infini_train/src/checkpoint/load_planner.cc @@ -1,146 +1,233 @@ #include "infini_train/include/checkpoint/load_planner.h" #include -#include -#include -#include +#include +#include #include "glog/logging.h" -#include "infini_train/include/checkpoint/checkpoint.h" -#include "infini_train/include/nn/parallel/global.h" - namespace infini_train::checkpoint { +namespace { -// ----------------------------------------------------------------------------- -// ReadItem helpers -// ----------------------------------------------------------------------------- - -static std::string DataTypeToString(DataType dt) { - auto it = kDataTypeToDesc.find(dt); - if (it != kDataTypeToDesc.end()) { - return it->second; +DataType StringToDataType(const std::string &value) { + for (const auto &[dtype, description] : kDataTypeToDesc) { + if (description == value) { + return dtype; + } } - return "fp32"; + return DataType::kFLOAT32; } -static DataType StringToDataType(const std::string &s) { - for (const auto &[dt, desc] : kDataTypeToDesc) { - if (desc == s) { - return dt; +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 DataType::kFLOAT32; + return fragmented_axis; } -bool ReadItem::IsLocalRead(int /*current_tp_rank*/, int /*current_tp_size*/, int /*current_pp_rank*/, - int /*current_pp_size*/) const { - - - return true; +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"; } -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- -std::vector LoadPlanner::Plan(const Checkpoint::CheckpointMetadata &metadata) { - std::vector items; - - for (const auto &entry : metadata.tensors) { - ReadItem item; - item.key = entry.key; - item.filename = entry.file; - item.dtype = StringToDataType(entry.dtype_str); - item.global_shape = entry.global_shape; - item.replica_id = 0; - item.byte_size = entry.byte_size; - - - - if (entry.key.find(".c_attn.") != std::string::npos || entry.key.find(".c_fc.") != std::string::npos - || entry.key.find(".c_fc2.") != std::string::npos || entry.key.find(".lm_head.") != std::string::npos - || entry.key.find(".wte.") != std::string::npos || entry.key.find(".output_layer.") != std::string::npos) { - ShardSpec spec; - spec.dim = 0; - spec.shard_count = metadata.parallel_config.tp_size; - spec.shard_index = 0; - spec.parallel_type = "tp"; - item.shard_specs.push_back(spec); - } else if (entry.key.find(".c_proj.") != std::string::npos) { - ShardSpec spec; - spec.dim = 1; - spec.shard_count = metadata.parallel_config.tp_size; - spec.shard_index = 0; - spec.parallel_type = "tp"; - item.shard_specs.push_back(spec); +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; } - - - items.push_back(std::move(item)); } - - return items; + return true; } -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- -std::unordered_map> -LoadPlanner::LoadFile(const std::filesystem::path &checkpoint_dir, const std::string &filename) { - auto path = checkpoint_dir / filename; - if (!std::filesystem::exists(path)) { - LOG(WARNING) << "[CKPT] File not found: " << path; - return {}; +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; } - return Checkpoint::LoadStateDictFile(path); } -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- -LoadPlanner::PlanResult LoadPlanner::PlanAndLoad(const std::filesystem::path &checkpoint_dir, - const Checkpoint::CheckpointMetadata &metadata) { - - PlanResult result; - - - result.read_items = Plan(metadata); - - - std::unordered_set filenames; - for (const auto &item : result.read_items) { filenames.insert(item.filename); } +} // 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; + } - for (const auto &filename : filenames) { - auto sd = LoadFile(checkpoint_dir, filename); - if (!sd.empty()) { - result.loaded_sd[filename] = std::move(sd); - LOG(INFO) << "[CKPT] Loaded " << filename << ": " << result.loaded_sd[filename].size() << " tensors"; + 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; + } - int missing = 0; - for (const auto &item : result.read_items) { - auto it = result.loaded_sd.find(item.filename); - if (it == result.loaded_sd.end()) { - LOG(WARNING) << "[CKPT] File " << item.filename << " not loaded, key=" << item.key; - missing++; + 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; } - if (it->second.find(item.key) == it->second.end()) { - LOG(WARNING) << "[CKPT] Key " << item.key << " not found in " << item.filename; - missing++; + + 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}); } - } - if (missing > 0) { - LOG(WARNING) << "[CKPT] " << missing << " tensors missing from checkpoint files"; + 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)); } - - LOG(INFO) << "[CKPT] LoadPlanner: " << result.read_items.size() << " read items, " << result.loaded_sd.size() - << " files loaded"; - - return result; + 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 00000000..47663788 --- /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 index c670a41f..1ef3e197 100644 --- a/infini_train/src/checkpoint/reshard.cc +++ b/infini_train/src/checkpoint/reshard.cc @@ -1,379 +1,71 @@ #include "infini_train/include/checkpoint/reshard.h" -#include #include #include -#include -#include -#include #include "glog/logging.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/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/process_group.h" -#include "infini_train/include/nn/parallel/tensor_parallel.h" -#include "infini_train/include/nn/parallel/utils.h" #include "infini_train/include/optimizer.h" -#include "infini_train/include/tensor.h" namespace infini_train::checkpoint { -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- - - -static std::shared_ptr SliceTensorDim(const std::shared_ptr &full, int dim, int64_t offset, - int64_t size) { - auto dims = full->Dims(); - CHECK_LT(dim, static_cast(dims.size())) - << "dim " << dim << " out of range for tensor with " << dims.size() << " dims"; - CHECK_GE(offset, 0) << "negative offset: " << offset; - CHECK_LE(offset + size, dims[dim]) << "slice [" << offset << ":" << offset + size << ") exceeds dimension " - << dims[dim]; - - auto sliced = full->Slice(dim, offset, offset + size); - return std::make_shared(sliced->To(full->GetDevice())); -} - - -static std::shared_ptr SliceForTPKey(const std::shared_ptr &full_tensor, const std::string &key, - int tp_rank, int tp_world_size) { - - auto dims = full_tensor->Dims(); - auto part_type = ReshardPlanner::InferPartitionType(key); - - if (part_type == ReshardPlanner::PartitionType::kColumnParallel) { - CHECK_GE(dims.size(), 1u) << "Expected at least 1D tensor for " << key; - auto [start, size] = GetRankSliceRange(dims[0], tp_world_size, tp_rank); - return SliceTensorDim(full_tensor, 0, start, size); - } - - if (part_type == ReshardPlanner::PartitionType::kRowParallel) { - CHECK_GE(dims.size(), 2u) << "Expected 2D tensor for " << key; - auto [start, size] = GetRankSliceRange(dims[1], tp_world_size, tp_rank); - return SliceTensorDim(full_tensor, 1, start, size); - } - - - return std::make_shared(*full_tensor); -} - -// ----------------------------------------------------------------------------- -// ReshardPlanner -// ----------------------------------------------------------------------------- - -ReshardPlanner::PartitionType ReshardPlanner::InferPartitionType(const std::string &key) { - - - if (key.find(".c_attn.") != std::string::npos || key.find(".c_fc.") != std::string::npos - || key.find(".c_fc2.") != std::string::npos || key.find(".lm_head.") != std::string::npos - || key.find(".wte.") != std::string::npos || key.find(".output_layer.") != std::string::npos - || key.find("linear_qkv.") != std::string::npos || key.find("linear_fc1_gate.") != std::string::npos - || key.find("linear_fc1_up.") != std::string::npos - || key.find("linear_fc1.weight") != std::string::npos) { // fused gate+up - return PartitionType::kColumnParallel; - } - - - - if (key.find(".c_proj.") != std::string::npos || key.find("linear_proj.") != std::string::npos - || key.find("linear_fc2.") != std::string::npos) { - return PartitionType::kRowParallel; - } - - - return PartitionType::kReplicated; -} - -ReshardPlan ReshardPlanner::ComputePlan(const std::vector &read_items, - const Checkpoint::CheckpointMetadata::ParallelConfig &saved_config, - const Checkpoint::CheckpointMetadata::ParallelConfig ¤t_config, - int tp_rank, int pp_rank) { - - ReshardPlan plan; - plan.saved_tp_size = saved_config.tp_size; - plan.saved_pp_size = saved_config.pp_size; - plan.current_tp_size = current_config.tp_size; - plan.current_pp_size = current_config.pp_size; - plan.tp_changed = saved_config.tp_size != current_config.tp_size; - plan.pp_changed = saved_config.pp_size != current_config.pp_size; - - for (const auto &item : read_items) { - ReshardTensorPlan tp; - tp.key = item.key; - - auto part_type = InferPartitionType(item.key); - tp.is_replicated = (part_type == PartitionType::kReplicated); - - if (item.global_shape.empty()) { - LOG(WARNING) << "[CKPT] No global_shape for key: " << item.key; - tp.needs_reshard = false; - tp.target_shape = {}; - plan.tensors.push_back(std::move(tp)); - continue; - } - - if (tp.is_replicated || !plan.tp_changed) { - - tp.target_shape = item.global_shape; - tp.saved_local_shape = item.global_shape; - tp.needs_reshard = false; - plan.tensors.push_back(std::move(tp)); - continue; - } - - - tp.needs_reshard = true; - - int shard_dim = (part_type == PartitionType::kColumnParallel) ? 0 : 1; - int current_world = current_config.tp_size; - int64_t global_size = item.global_shape[shard_dim]; - - - auto [start, size] = GetRankSliceRange(global_size, current_world, tp_rank); - - - tp.target_shape = item.global_shape; - tp.target_shape[shard_dim] = size; - tp.saved_local_shape = item.global_shape; - - - ReshardSlice slice; - slice.dim = shard_dim; - slice.global_offset = start; - slice.size = size; - tp.slices.push_back(slice); - - plan.tensors.push_back(std::move(tp)); - } - - return plan; -} - -// ----------------------------------------------------------------------------- -// ReshardExecutor -// ----------------------------------------------------------------------------- - -std::unordered_map> -ReshardExecutor::Execute(const ReshardPlan &plan, - const std::unordered_map> &full_sd, int tp_rank, - int pp_rank) { - - (void)pp_rank; - - std::unordered_map> result; - int tp_world_size = nn::parallel::global::GetTensorParallelSize(); - - for (const auto &tp : plan.tensors) { - auto it = full_sd.find(tp.key); - if (it == full_sd.end()) { - LOG(WARNING) << "[CKPT] Reshard: key not found in full state_dict: " << tp.key; - continue; - } - - if (!tp.needs_reshard) { - - result[tp.key] = it->second; - continue; - } - - - auto tensor = it->second; - - if (!tp.slices.empty()) { - if (tp.slices.size() == 1) { - - const auto &slice = tp.slices[0]; - result[tp.key] = SliceTensorDim(tensor, slice.dim, slice.global_offset, slice.size); - } else { - - auto sliced = tensor; - for (const auto &slice : tp.slices) { - sliced = SliceTensorDim(sliced, slice.dim, slice.global_offset, slice.size); - } - result[tp.key] = sliced; - } - } else { - - result[tp.key] = SliceForTPKey(tensor, tp.key, tp_rank, tp_world_size); - } - } - - - for (const auto &[key, tensor] : full_sd) { - if (result.find(key) == result.end()) { - result[key] = tensor; - } - } - - LOG(INFO) << "[CKPT] ReshardExecutor: " << result.size() << " tensors ready for rank"; - - return result; -} - -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- - -static const nn::parallel::ProcessGroup *GetTPGroup() { - auto device = Device(); - auto tp_group_name = nn::parallel::GetTensorParallelProcessGroupName(device.Rank().GlobalRank()); - return nn::parallel::ProcessGroupFactory::Instance()->Get(tp_group_name); -} - - - - - -std::unordered_map> -AllGatherFullModel(const std::unordered_map> &local_sd, int saved_tp_size) { - - auto tp_group = GetTPGroup(); - CHECK(tp_group != nullptr) << "TP process group not found, cannot gather for resharding"; - - std::unordered_map> full_sd; - - for (const auto &[key, local_tensor] : local_sd) { - auto local_dims = local_tensor->Dims(); - auto part_type = ReshardPlanner::InferPartitionType(key); - - - int shard_dim = (part_type == ReshardPlanner::PartitionType::kRowParallel) ? 1 : 0; - - if (shard_dim == 0) { - - std::vector full_dims = local_dims; - full_dims[0] *= saved_tp_size; - - auto full_tensor = std::make_shared(full_dims, local_tensor->Dtype(), local_tensor->GetDevice()); - tp_group->AllGather(full_tensor, local_tensor, false); - full_sd[key] = full_tensor; - - } else { - - - - auto local_transposed = local_tensor->Transpose(0, 1); - - - std::vector gathered_dims = local_transposed->Dims(); - gathered_dims[0] *= saved_tp_size; - - auto gathered = std::make_shared(gathered_dims, local_tensor->Dtype(), local_tensor->GetDevice()); - tp_group->AllGather(gathered, local_transposed, false); - - - - auto full_tensor = std::make_shared(gathered->Transpose(0, 1)->To(local_tensor->GetDevice())); - full_sd[key] = full_tensor; - } - } - - return full_sd; -} - -// ----------------------------------------------------------------------------- - -// ----------------------------------------------------------------------------- - -void ReshardAndLoad(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, - TrainerState &state, bool load_optimizer_state, LRScheduler *lr_scheduler, - const Checkpoint::CheckpointMetadata &metadata, int tp_rank, int pp_rank, int sp_rank) { - - (void)sp_rank; - - int saved_tp = metadata.parallel_config.tp_size; - int current_tp = nn::parallel::global::GetTensorParallelSize(); - bool tp_changed = saved_tp != current_tp; - - - auto load_result = LoadPlanner::PlanAndLoad(checkpoint_dir, metadata); - auto &read_items = load_result.read_items; - auto &loaded_sd = load_result.loaded_sd; - - - auto local_sd_it = loaded_sd.find("model.ckpt"); - CHECK(local_sd_it != loaded_sd.end()) << "model.ckpt not loaded"; - auto local_sd = std::move(local_sd_it->second); - - LOG(INFO) << "[CKPT] Loaded local shard: " << local_sd.size() << " tensors"; - - std::unordered_map> target_sd; - - if (tp_changed) { - - LOG(INFO) << "[CKPT] TP changed: " << saved_tp << " -> " << current_tp - << ". Gathering full model for resharding..."; - - auto full_sd = AllGatherFullModel(local_sd, saved_tp); - LOG(INFO) << "[CKPT] Gathered full model: " << full_sd.size() << " tensors"; - - - auto reshard_plan - = ReshardPlanner::ComputePlan(read_items, metadata.parallel_config, - Checkpoint::CheckpointMetadata::ParallelConfig{ - current_tp, metadata.parallel_config.pp_size, - metadata.parallel_config.dp_size, metadata.parallel_config.sp_size}, - tp_rank, pp_rank); - - - target_sd = ReshardExecutor::Execute(reshard_plan, full_sd, tp_rank, pp_rank); - - LOG(INFO) << "[CKPT] Resharded to TP=" << current_tp << ", rank " << tp_rank << " has " << target_sd.size() - << " tensors"; - } else { - - target_sd = std::move(local_sd); - } - - - model.LoadStateDict(target_sd); - - - if (load_optimizer_state && optimizer != nullptr) { - auto opt_it = loaded_sd.find("optimizer.ckpt"); - if (opt_it != loaded_sd.end()) { - if (tp_changed) { - - auto opt_t = opt_it->second.find("adam.t"); - if (opt_t != opt_it->second.end()) { - optimizer->LoadStateDict({{"adam.t", opt_t->second}}); - LOG(WARNING) << "[CKPT] TP changed, optimizer m/v states reinitialized " - << "(only step counter loaded)"; - } - } else { - optimizer->LoadStateDict(opt_it->second); - } - } - } - +void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, + TrainerState &state, bool load_optimizer_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; - - state.tp_size = saved_tp; - state.pp_size = metadata.parallel_config.pp_size; - state.sp_size = metadata.parallel_config.sp_size; - state.ddp_size = metadata.parallel_config.dp_size; - - - if (lr_scheduler != nullptr) { - auto lr_path = checkpoint_dir / "lr_scheduler.ckpt"; - if (std::filesystem::exists(lr_path)) { - lr_scheduler->LoadStateDict(Checkpoint::LoadLRSchedulerStateFile(lr_path)); + if (load_optimizer_state && 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 { - LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found: " << lr_path; + 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)); } } - - LOG(ERROR) << "[CKPT] ReshardAndLoad done: global_step=" << state.global_step - << ", consumed_batches=" << state.consumed_batches << ", TP=" << current_tp - << ", tensors=" << target_sd.size(); + 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 284d3d0c..cf2a3926 100644 --- a/infini_train/src/checkpoint/save_planner.cc +++ b/infini_train/src/checkpoint/save_planner.cc @@ -1,7 +1,50 @@ #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; @@ -18,7 +61,8 @@ std::vector SavePlanner::Plan(const ShardedStateDict &sd, int rank) { item.byte_size = TensorByteSize(info.dtype, info.local_shape); item.dtype = info.dtype; item.local_shape = info.local_shape; - item.shard_specs = info.shard_specs; + item.global_offset = info.global_offset; + item.axis_fragmentations = info.axis_fragmentations; item.replica_id = info.replica_id; item.rank = rank; diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index f4ba274c..0bf4b49e 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -55,25 +55,42 @@ Module::NamedParameters(const std::string &prefix, bool recurse, bool remove_dup std::function collect = [&](const Module &module, const std::string &module_prefix) { + std::vector>> parameters; + parameters.reserve(module.parameters_.size()); for (const auto &[name, parameter] : module.parameters_) { - if (!parameter) { - continue; + if (parameter) { + parameters.emplace_back(name, parameter); } + } + std::sort(parameters.begin(), parameters.end(), + [](const auto &left, const auto &right) { return left.first < right.first; }); + + for (auto &[name, parameter] : parameters) { if (remove_duplicate && !visited.insert(parameter.get()).second) { continue; } const auto full_name = module_prefix.empty() ? name : module_prefix + "." + name; - named_parameters.emplace_back(full_name, parameter); + named_parameters.emplace_back(full_name, std::move(parameter)); } if (!recurse) { return; } + std::vector>> children; + children.reserve(module.modules_.size()); for (const auto &[name, child] : module.modules_) { - if (!child) { + if (name.starts_with("__pp")) { continue; } + if (child) { + children.emplace_back(name, child); + } + } + std::sort(children.begin(), children.end(), + [](const auto &left, const auto &right) { return left.first < right.first; }); + + for (const auto &[name, child] : children) { const auto child_prefix = module_prefix.empty() ? name : module_prefix + "." + name; collect(*child, child_prefix); } @@ -182,28 +199,29 @@ std::unordered_map> Module::StateDict() con return state; } -checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix, - const std::vector &parent_shards) const { +checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix) const { checkpoint::ShardedStateDict sd; for (auto &[name, param] : parameters_) { - checkpoint::ShardedTensorInfo info; + 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.shard_specs = parent_shards; + info.global_offset.assign(param->Dims().size(), 0); + info.axis_fragmentations.assign(param->Dims().size(), 1); info.replica_id = 0; sd.tensors[info.key] = std::move(info); } for (auto &[name, buffer] : buffers_) { - checkpoint::ShardedTensorInfo info; + 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.shard_specs = parent_shards; + info.global_offset.assign(buffer->Dims().size(), 0); + info.axis_fragmentations.assign(buffer->Dims().size(), 1); info.replica_id = 0; sd.tensors[info.key] = std::move(info); } @@ -214,7 +232,7 @@ checkpoint::ShardedStateDict Module::ShardedStateDict(const std::string &prefix, } auto child_prefix = prefix.empty() ? name : prefix + "." + name; - auto child_sd = module->ShardedStateDict(child_prefix, parent_shards); + auto child_sd = module->ShardedStateDict(child_prefix); sd.Merge(std::move(child_sd)); } 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 bc4c6bd4..1385e9ba 100644 --- a/infini_train/src/nn/modules/transformer/causal_self_attention.cc +++ b/infini_train/src/nn/modules/transformer/causal_self_attention.cc @@ -76,6 +76,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 99a739d2..d6e6e28a 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 b08b64fb..8cd51351 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -201,6 +201,11 @@ 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::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 c0369cde..0dc34fdf 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 cf858340..5f3caf6d 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -284,40 +284,31 @@ 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 std::vector &parent_shards) const { +checkpoint::ShardedStateDict ColumnParallelLinear::ShardedStateDict(const std::string &prefix) const { checkpoint::ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); - // Weight is split along output dimension (dim=0) - auto tp_shards = parent_shards; - tp_shards.push_back({ - .dim = 0, - .shard_count = tp_size, - .shard_index = tp_rank, - .parallel_type = "tp", - }); - auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensorInfo w; - w.key = prefix + "." + 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.shard_specs = tp_shards; + w.global_offset = {output_size_per_partition_ * tp_rank, 0}; + w.axis_fragmentations = {tp_size, 1}; w.replica_id = 0; sd.tensors[w.key] = std::move(w); // Bias is also split along dim=0 if (bias_) { auto &bias = parameter(kParamBiasName); - checkpoint::ShardedTensorInfo b; - b.key = prefix + "." + 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.shard_specs = tp_shards; + b.global_offset = {output_size_per_partition_ * tp_rank}; + b.axis_fragmentations = {tp_size}; b.replica_id = 0; sd.tensors[b.key] = std::move(b); } @@ -380,40 +371,31 @@ 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 std::vector &parent_shards) const { +checkpoint::ShardedStateDict RowParallelLinear::ShardedStateDict(const std::string &prefix) const { checkpoint::ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); - // Weight is split along input dimension (dim=1) - auto tp_shards = parent_shards; - tp_shards.push_back({ - .dim = 1, - .shard_count = tp_size, - .shard_index = tp_rank, - .parallel_type = "tp", - }); - auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensorInfo w; - w.key = prefix + "." + 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.shard_specs = tp_shards; + w.global_offset = {0, input_size_per_partition_ * tp_rank}; + w.axis_fragmentations = {1, tp_size}; w.replica_id = 0; sd.tensors[w.key] = std::move(w); // Bias is NOT sharded in RowParallelLinear if (bias_) { auto &bias = parameter(kParamBiasName); - checkpoint::ShardedTensorInfo b; - b.key = prefix + "." + 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.shard_specs = parent_shards; // no TP shard for bias + b.global_offset = {0}; + b.axis_fragmentations = {1}; b.replica_id = 0; sd.tensors[b.key] = std::move(b); } @@ -478,27 +460,18 @@ VocabParallelEmbedding::Forward(const std::vector> &inpu return {output}; } -checkpoint::ShardedStateDict -VocabParallelEmbedding::ShardedStateDict(const std::string &prefix, - const std::vector &parent_shards) const { +checkpoint::ShardedStateDict VocabParallelEmbedding::ShardedStateDict(const std::string &prefix) const { checkpoint::ShardedStateDict sd; int tp_size = global::GetTensorParallelSize(); - auto tp_shards = parent_shards; - tp_shards.push_back({ - .dim = 0, - .shard_count = tp_size, - .shard_index = tp_rank, - .parallel_type = "tp", - }); - auto &weight = parameter(kParamWeightName); - checkpoint::ShardedTensorInfo w; - w.key = prefix + "." + 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.shard_specs = tp_shards; + w.global_offset = {vocab_start_index_, 0}; + w.axis_fragmentations = {tp_size, 1}; w.replica_id = 0; sd.tensors[w.key] = std::move(w); @@ -575,7 +548,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). auto softmax_local = exp_local->Div(sum_exp->Unsqueeze(-1)); // 5. Perform allreduce to get global predicted_logit @@ -589,7 +562,7 @@ VocabParallelCrossEntropy::Forward(const std::vector> &i auto log_sum_exp = sum_exp->Log(); auto loss = log_sum_exp->Sub(predicted); - + // 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 11c1d4ef..9a158061 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,6 +18,47 @@ 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_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { @@ -54,4 +99,341 @@ 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, 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 1cbb8b9f..e3b0ebdf 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 d4a6efc2..b6fc6111 100644 --- a/tests/transformer/test_transformer_architecture.cc +++ b/tests/transformer/test_transformer_architecture.cc @@ -159,6 +159,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 a3f618125a5a64b24ddd030d88c6538c84fe400b Mon Sep 17 00:00:00 2001 From: JYMiracle305 <604951424@qq.com> Date: Mon, 3 Aug 2026 02:53:44 +0000 Subject: [PATCH 3/3] feat: add checkpoint validation and save synchronization --- infini_train/include/checkpoint/shard_spec.h | 8 ++++- .../src/checkpoint/checkpoint_manager.cc | 34 ++++++++++++++++--- infini_train/src/checkpoint/load_planner.cc | 11 ++++++ .../test_checkpoint_serialization.cc | 18 ++++++++++ 4 files changed, 66 insertions(+), 5 deletions(-) diff --git a/infini_train/include/checkpoint/shard_spec.h b/infini_train/include/checkpoint/shard_spec.h index d6c48951..dc52ddb4 100644 --- a/infini_train/include/checkpoint/shard_spec.h +++ b/infini_train/include/checkpoint/shard_spec.h @@ -5,6 +5,8 @@ #include #include +#include "glog/logging.h" + #include "infini_train/include/datatype.h" namespace infini_train::checkpoint { @@ -44,7 +46,11 @@ struct ShardedStateDict { std::map tensors; void Merge(ShardedStateDict &&other) { - for (auto &[key, info] : other.tensors) { tensors.emplace(std::move(key), std::move(info)); } + 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; + } } }; diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index 5b2faddc..4063c7f5 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -17,6 +17,9 @@ #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; @@ -36,14 +39,29 @@ std::filesystem::path ResolveCheckpointDirectory(const std::filesystem::path &ro return directory; } -void WaitForWriterManifests(const std::filesystem::path &staging_root, int tp_size, int pp_size) { +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); - if (!std::filesystem::exists(staging_root / std::format("rank_{:06d}", rank) / "metadata.json")) { + 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; } @@ -88,7 +106,10 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & 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_micro_batches = static_cast(std::max(args.state.consumed_micro_batches, 0)); if (args.rank.IsMainRank()) { @@ -116,6 +137,12 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { : 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) { @@ -135,7 +162,6 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { Checkpoint::SaveSharded(rank_dir, sharded_state, write_items, args.model.StateDict(), optimizer_state, state, args.rank.GlobalRank()); - const auto staging_root = iteration_dir / ".metadata_tmp"; 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"; @@ -149,7 +175,7 @@ void SaveCheckpoint(const SaveCheckpointArgs &args) { 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); + 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"; diff --git a/infini_train/src/checkpoint/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc index 334268f4..871070d1 100644 --- a/infini_train/src/checkpoint/load_planner.cc +++ b/infini_train/src/checkpoint/load_planner.cc @@ -3,6 +3,7 @@ #include #include #include +#include #include "glog/logging.h" @@ -10,11 +11,21 @@ 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; } diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 9a158061..fe47c631 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -61,6 +61,15 @@ TEST(ModuleNamedParametersTest, SupportsTorchStyleArgumentsAndSharedParameterDed 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); @@ -283,6 +292,15 @@ TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { 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"),