diff --git a/example/gpt2/checkpoint_loader.cc b/example/gpt2/checkpoint_loader.cc index 95e54730b..ceca604be 100644 --- a/example/gpt2/checkpoint_loader.cc +++ b/example/gpt2/checkpoint_loader.cc @@ -163,15 +163,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight - int local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { - auto &tensor - = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), - nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)]; + auto &tensor = state_dict[std::format( + "{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, + std::to_string(idx), nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_w_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_w_bytes, std::ios::cur); @@ -179,14 +176,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.bias - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_b_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_b_bytes, std::ios::cur); @@ -194,13 +189,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_attn.weight (ColumnParallelLinear, but actually applies on "rows") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; // NOTE(zbl): In the .bin model file, Q/K/V is concated along last dim, // i.e. [Q|K|V].T = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn].T // However, each tp_rank needs to get [q_i|k_i|v_i].T, so we need to jump and read them @@ -229,7 +223,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/rows_all, /*cols=*/cols_all, /*row_start=*/2 * n_embd + tp_rank * local_C, /*row_cnt=*/local_C); - ++local_layer_index; } else { size_t c_attn_w_bytes = qkv_out * n_embd * sizeof(float); ifs.seekg(c_attn_w_bytes, std::ios::cur); @@ -237,13 +230,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_attn.bias (ColumnParallelLinear) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamBiasName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamBiasName)]; // NOTE(zbl): Same as c_attn.weight, the bias for Q/K/V is concated // i.e. [Q|K|V] = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn] // However, each tp_rank needs to get [q_i|k_i|v_i], so we need to jump and read them @@ -271,7 +263,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*len=*/len_all, /*start=*/2 * n_embd + tp_rank * local_C, /*cnt=*/local_C); - ++local_layer_index; } else { size_t c_attn_b_bytes = qkv_out * sizeof(float); ifs.seekg(c_attn_b_bytes, std::ios::cur); @@ -279,16 +270,14 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight (RowParallelLinear, but actually applies on "columns") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), n_embd, n_embd, tp_rank * in_pp, in_pp); - ++local_layer_index; } else { size_t c_proj_w_bytes = n_embd * n_embd * sizeof(float); ifs.seekg(c_proj_w_bytes, std::ios::cur); @@ -296,15 +285,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.bias (RowParallelLinear, no shard on bias) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamBiasName)]; + std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t c_proj_b_bytes = n_embd * sizeof(float); ifs.seekg(c_proj_b_bytes, std::ios::cur); @@ -312,15 +299,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { - auto &tensor - = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), - nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)]; + auto &tensor = state_dict[std::format( + "{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, + std::to_string(idx), nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_w_bytes = n_embd * sizeof(float); ifs.seekg(ln_2_w_bytes, std::ios::cur); @@ -328,14 +312,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.bias - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_b_bytes = n_embd * sizeof(float); ifs.seekg(ln_2_b_bytes, std::ios::cur); @@ -343,15 +325,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (ColumnParallelLinear, but actually applies on "rows") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), fc_out, n_embd, fc_start, fc_pp); - ++local_layer_index; } else { size_t c_fc_w_bytes = fc_out * n_embd * sizeof(float); ifs.seekg(c_fc_w_bytes, std::ios::cur); @@ -359,15 +339,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.bias (ColumnParallelLinear) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamBiasName)]; ReadVectorShardFloat(ifs, static_cast(tensor->DataPtr()), fc_out, fc_start, fc_pp); - ++local_layer_index; } else { size_t c_fc_b_bytes = fc_out * sizeof(float); ifs.seekg(c_fc_b_bytes, std::ios::cur); @@ -375,16 +353,14 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight (RowParallelLinear, but actually applies on "columns") - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), n_embd, fc_out, tp_rank * in4_pp, in4_pp); - ++local_layer_index; } else { size_t c_proj_w_bytes = fc_out * n_embd * sizeof(float); ifs.seekg(c_proj_w_bytes, std::ios::cur); @@ -392,15 +368,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.bias (RowParallelLinear, no shard on bias) - local_layer_index = 0; for (int idx = 0; idx < n_layer; ++idx) { if (owned_layers[idx]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(idx), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamBiasName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t c_proj_b_bytes = n_embd * sizeof(float); ifs.seekg(c_proj_b_bytes, std::ios::cur); diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 4af5b8260..f63c808d6 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -421,11 +421,13 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, diff --git a/example/llama3/checkpoint_loader.cc b/example/llama3/checkpoint_loader.cc index 32b11e430..5622afc44 100644 --- a/example/llama3/checkpoint_loader.cc +++ b/example/llama3/checkpoint_loader.cc @@ -181,14 +181,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight : Full version nn::RMSNorm - int local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn1LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_bytes, std::ios::cur); @@ -197,13 +195,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) // transformer.h.{i}.attn.c_attn.weight : ColumnParallelLinear, but actually applies on "rows" // W-qkv should be [Q(=n_embd) | K(=n_kv_head*head_dim) | V(=n_kv_head*head_dim)] × n_embd - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()); const std::streampos base_pos = ifs.tellg(); @@ -229,7 +226,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/attn_rows_all, /*cols=*/attn_cols, /*row_start=*/q_out_rows + kv_out_rows + tp_rank * kv_local_rows, /*row_cnt=*/kv_local_rows); - ++local_layer_index; } else { size_t qkv_bytes = static_cast(attn_rows_all) * attn_cols * sizeof(float); ifs.seekg(qkv_bytes, std::ios::cur); @@ -237,17 +233,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/n_embd, /*col_start=*/tp_rank * in_pp, /*col_cnt=*/in_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * n_embd * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); @@ -255,14 +249,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight : Full version RMSNorm - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn2LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_bytes = static_cast(n_embd) * sizeof(float); ifs.seekg(ln_2_bytes, std::ios::cur); @@ -270,18 +262,16 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (up) -> local packed c_fc rows [fc_pp : 2*fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()) + fc_pp * n_embd; ReadMatrixRowShardFloat(ifs, dst, /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc_bytes, std::ios::cur); @@ -289,17 +279,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc2.weight (gate) -> local packed c_fc rows [0 : fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc2_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc2_bytes, std::ios::cur); @@ -307,17 +295,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/fc_out, /*col_start=*/tp_rank * in_fc_pp, /*col_cnt=*/in_fc_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * ffn_hidden * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); diff --git a/example/llama3/main.cc b/example/llama3/main.cc index b8342ca90..46d8c8a64 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -411,11 +411,13 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, diff --git a/example/qwen3/checkpoint_loader.cc b/example/qwen3/checkpoint_loader.cc index 181217a32..7b9b9e65e 100644 --- a/example/qwen3/checkpoint_loader.cc +++ b/example/qwen3/checkpoint_loader.cc @@ -201,35 +201,31 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_1.weight : Full version nn::RMSNorm - int local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn1LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_1_bytes = n_embd * sizeof(float); ifs.seekg(ln_1_bytes, std::ios::cur); } } - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &q_norm_tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kQNormLayerName, nn::RMSNorm::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kQNormLayerName, + nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(q_norm_tensor->DataPtr()), head_dim); auto &k_norm_tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kKNormLayerName, nn::RMSNorm::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kKNormLayerName, + nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(k_norm_tensor->DataPtr()), head_dim); - ++local_layer_index; } else { size_t qk_norm_bytes = 2 * head_dim * sizeof(float); ifs.seekg(qk_norm_bytes, std::ios::cur); @@ -238,13 +234,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) // transformer.h.{i}.attn.c_attn.weight : ColumnParallelLinear, but actually applies on "rows" // W-qkv should be [Q(=n_embd) | K(=n_kv_head*head_dim) | V(=n_kv_head*head_dim)] x n_embd - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName, + nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()); const std::streampos base_pos = ifs.tellg(); @@ -270,7 +265,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) /*rows=*/attn_rows_all, /*cols=*/attn_cols, /*row_start=*/q_out_rows + kv_out_rows + tp_rank * kv_local_rows, /*row_cnt=*/kv_local_rows); - ++local_layer_index; } else { size_t qkv_bytes = static_cast(attn_rows_all) * attn_cols * sizeof(float); ifs.seekg(qkv_bytes, std::ios::cur); @@ -278,17 +272,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, - std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, - nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; + std::to_string(i), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName, + nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/n_embd, /*col_start=*/tp_rank * in_pp, /*col_cnt=*/in_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * n_embd * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); @@ -296,14 +288,12 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.ln_2.weight : Full version RMSNorm - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kLn2LayerName, nn::RMSNorm::kParamWeightName)]; ReadVectorAllFloat(ifs, static_cast(tensor->DataPtr()), n_embd); - ++local_layer_index; } else { size_t ln_2_bytes = static_cast(n_embd) * sizeof(float); ifs.seekg(ln_2_bytes, std::ios::cur); @@ -311,17 +301,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc2.weight (gate) -> local packed c_fc rows [0 : fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc_bytes, std::ios::cur); @@ -329,17 +317,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_fc.weight (up) -> local packed c_fc rows [fc_pp : 2*fc_pp) - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; float *dst = static_cast(tensor->DataPtr()) + fc_pp * n_embd; ReadMatrixRowShardFloat(ifs, dst, /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); - ++local_layer_index; } else { size_t fc2_bytes = static_cast(ffn_hidden) * n_embd * sizeof(float); ifs.seekg(fc2_bytes, std::ios::cur); @@ -347,17 +333,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; for (int i = 0; i < static_cast(n_layer); ++i) { if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, - nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), + nn::TransformerChunk::kHLayerName, std::to_string(i), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)]; ReadMatrixColShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/n_embd, /*cols=*/fc_out, /*col_start=*/tp_rank * in_fc_pp, /*col_cnt=*/in_fc_pp); - ++local_layer_index; } else { size_t c_proj_bytes = static_cast(n_embd) * ffn_hidden * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); diff --git a/example/qwen3/main.cc b/example/qwen3/main.cc index 11249f9fc..763afa1e7 100644 --- a/example/qwen3/main.cc +++ b/example/qwen3/main.cc @@ -411,11 +411,13 @@ void Train(const nn::parallel::Rank &rank) { .n_head = model_config.n_head, .n_kv_head = model_config.n_kv_head, .n_embd = model_config.n_embd, - .vocab_size = model_config.vocab_size, + .original_vocab_size = model_config.original_vocab_size, + .padded_vocab_size = model_config.vocab_size, .ddp_size = ddp_world_size, .tp_size = tp_world_size, .sp_size = sp_world_size, .pp_size = pp_world_size, + .vpp_size = static_cast(FLAGS_virtual_pipeline_parallel), .checkpoint_root_dir = FLAGS_save, .max_checkpoint_keep = FLAGS_max_checkpoint_keep, .rank = rank, diff --git a/infini_train/include/checkpoint/checkpoint.h b/infini_train/include/checkpoint/checkpoint.h index a69aad232..9703cdc7b 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/lr_scheduler.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Optimizer; @@ -22,11 +27,13 @@ struct TrainerState { int64_t n_head = 0; int64_t n_kv_head = 0; int64_t n_embd = 0; - int64_t vocab_size = 0; + int64_t original_vocab_size = 0; + int64_t padded_vocab_size = 0; int ddp_size = 1; int tp_size = 1; int sp_size = 1; int pp_size = 1; + int vpp_size = 1; }; class Checkpoint { @@ -37,12 +44,56 @@ class Checkpoint { static void Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, LRScheduler *lr_scheduler); + 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; + + struct TensorEntry { + std::string key; + std::string dtype_str; + std::vector global_shape; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + std::vector segments; + std::string file; + uint64_t offset = 0; + uint64_t byte_size = 0; + std::vector stored_on_ranks; + int pp_rank = 0; + }; + + std::vector tensors; + bool has_metadata = false; + }; + + static CheckpointMetadata LoadMetadata(const std::filesystem::path &checkpoint_dir); + static void SaveMetadataFile(const std::filesystem::path &path, const CheckpointMetadata &metadata); + 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); + static void SaveLocalShard(const std::filesystem::path &checkpoint_dir, const ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + int global_rank); + static void SaveTrainerState(const std::filesystem::path &path, const TrainerState &state); static TrainerState LoadTrainerState(const std::filesystem::path &path); }; diff --git a/infini_train/include/checkpoint/checkpoint_manager.h b/infini_train/include/checkpoint/checkpoint_manager.h index 13490ccce..fa4a3cc61 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" @@ -45,11 +44,13 @@ struct SaveCheckpointArgs { int64_t n_head = 0; int64_t n_kv_head = 0; int64_t n_embd = 0; - int64_t vocab_size = 0; + int64_t original_vocab_size = 0; + int64_t padded_vocab_size = 0; int ddp_size = 1; int tp_size = 1; int sp_size = 1; int pp_size = 1; + int vpp_size = 1; std::filesystem::path checkpoint_root_dir; size_t max_checkpoint_keep = 0; const nn::parallel::Rank &rank; diff --git a/infini_train/include/checkpoint/constants.h b/infini_train/include/checkpoint/constants.h new file mode 100644 index 000000000..e7ad0869c --- /dev/null +++ b/infini_train/include/checkpoint/constants.h @@ -0,0 +1,14 @@ +#pragma once + +namespace infini_train::checkpoint { + +inline constexpr char kModelCheckpointFilename[] = "model.ckpt"; +inline constexpr char kOptimizerCheckpointFilename[] = "optimizer.ckpt"; +inline constexpr char kMetadataFilename[] = "metadata.json"; +inline constexpr char kTemporaryMetadataFilename[] = "metadata.json.tmp"; +inline constexpr char kTrainerStateFilename[] = "trainer_state.json"; +inline constexpr char kLRSchedulerFilename[] = "lr_scheduler.ckpt"; +inline constexpr char kLatestIterationFilename[] = "latest_checkpointed_iteration.txt"; +inline constexpr char kTemporaryLatestIterationFilename[] = "latest_checkpointed_iteration.txt.tmp"; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/load_planner.h b/infini_train/include/checkpoint/load_planner.h new file mode 100644 index 000000000..bbca3fc75 --- /dev/null +++ b/infini_train/include/checkpoint/load_planner.h @@ -0,0 +1,52 @@ +#pragma once + +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/datatype.h" +#include "infini_train/include/shard_spec.h" + +namespace infini_train::checkpoint { + +// One storage-region transfer from a saved shard into a target local tensor. +struct ReadItem { + std::string key; + std::string filename; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + uint64_t byte_size = 0; + uint64_t data_offset = 0; + int shard_dim = -1; + int64_t source_offset = 0; + int64_t target_offset = 0; + int64_t length = 0; + std::vector source_shape; +}; + +// All reads required to materialize one target local tensor. +struct TargetTensorPlan { + std::string key; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector target_shape; + int shard_dim = -1; + int64_t trailing_zero_fill = 0; + std::vector reads; +}; + +// Complete load plan for one rank. +struct LoadPlan { + std::map tensors; +}; + +class LoadPlanner { +public: + // Compute saved-to-target overlaps from explicit global shard coordinates. + static LoadPlan PlanReshard(const Checkpoint::CheckpointMetadata &metadata, + const ShardedStateDict &target_state_dict); +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/load_strategy.h b/infini_train/include/checkpoint/load_strategy.h new file mode 100644 index 000000000..f5f6e2c0f --- /dev/null +++ b/infini_train/include/checkpoint/load_strategy.h @@ -0,0 +1,30 @@ +#pragma once + +#include +#include +#include +#include + +#include "infini_train/include/checkpoint/load_planner.h" + +namespace infini_train { +class Tensor; +} + +namespace infini_train::checkpoint { + +using LoadedStateDict = std::unordered_map>; + +class LoadStrategy { +public: + virtual ~LoadStrategy() = default; + virtual LoadedStateDict Execute(const std::filesystem::path &checkpoint_dir, const LoadPlan &plan) = 0; +}; + +/// Reads source regions directly from metadata offsets while caching one open stream per file. +class IndexedRegionLoadStrategy final : public LoadStrategy { +public: + LoadedStateDict Execute(const std::filesystem::path &checkpoint_dir, const LoadPlan &plan) override; +}; + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/checkpoint/save_planner.h b/infini_train/include/checkpoint/save_planner.h new file mode 100644 index 000000000..3c5a6fe3c --- /dev/null +++ b/infini_train/include/checkpoint/save_planner.h @@ -0,0 +1,47 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include "infini_train/include/datatype.h" +#include "infini_train/include/shard_spec.h" + +namespace infini_train { +class Tensor; +} + +namespace infini_train::checkpoint { + +// Physical write description for one local tensor shard. +struct WriteItem { + std::string key; + std::string filename; // Checkpoint payload filename. + uint64_t byte_size = 0; // Tensor payload size in bytes. + DataType dtype = DataType::kFLOAT32; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + int rank = 0; +}; + +// Build the local tensor write layout from a ShardedStateDict. +class SavePlanner { +public: + static std::vector Plan(const ShardedStateDict &sd, int rank); +}; + +ShardedStateDict +BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, + const std::unordered_map> &optimizer_state); + +// Return the number of payload bytes required by a tensor. +inline uint64_t TensorByteSize(DataType dtype, const std::vector &shape) { + uint64_t numel = 1; + for (auto d : shape) { numel *= static_cast(d); } + return numel * static_cast(kDataTypeToSize.at(dtype)); +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/include/nn/lora/lora_parallel_linear.h b/infini_train/include/nn/lora/lora_parallel_linear.h index d73a6e2b5..5c9bbddba 100644 --- a/infini_train/include/nn/lora/lora_parallel_linear.h +++ b/infini_train/include/nn/lora/lora_parallel_linear.h @@ -34,6 +34,8 @@ class LoRAColumnParallelLinear : public nn::parallel::ColumnParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + void MergeWeights(); void UnmergeWeights(); bool IsMerged() const; @@ -74,6 +76,8 @@ class LoRARowParallelLinear : public nn::parallel::RowParallelLinear { std::vector> Forward(const std::vector> &input_tensors) override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + void MergeWeights(); void UnmergeWeights(); bool IsMerged() const; diff --git a/infini_train/include/nn/modules/module.h b/infini_train/include/nn/modules/module.h index 8570c4768..40c5f83ac 100644 --- a/infini_train/include/nn/modules/module.h +++ b/infini_train/include/nn/modules/module.h @@ -8,6 +8,7 @@ #include "infini_train/include/datatype.h" #include "infini_train/include/device.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Tensor; @@ -63,7 +64,7 @@ class Module : public std::enable_shared_from_this { // InfiniTrain's NamedParameters returns results ordered by full parameter name. // TODO: Align with PyTorch's ordering in the future. - std::vector>> + virtual std::vector>> NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const; bool has_parameter(const std::string &name) const; std::shared_ptr *mutable_parameter(const std::string &name); @@ -75,11 +76,14 @@ class Module : public std::enable_shared_from_this { std::shared_ptr &mutable_module(const std::string &name); const Module &module(const std::string &name) const; - std::unordered_map> StateDict() const; + virtual std::unordered_map> StateDict() const; + + // Return state-dict metadata with global shard coordinates. + virtual ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const; // Current behavior: missing keys / shape / dtype mismatches are FATAL errors; unexpected keys in state_dict are // WARNING-only and silently ignored. - void LoadStateDict(const std::unordered_map> &state_dict); + virtual void LoadStateDict(const std::unordered_map> &state_dict); // operator() calls hooks and Forward std::vector> operator()(const std::vector> &input_tensors); diff --git a/infini_train/include/nn/modules/transformer/causal_self_attention.h b/infini_train/include/nn/modules/transformer/causal_self_attention.h index a3fa6d257..15f2f8b1e 100644 --- a/infini_train/include/nn/modules/transformer/causal_self_attention.h +++ b/infini_train/include/nn/modules/transformer/causal_self_attention.h @@ -25,6 +25,8 @@ class CausalSelfAttention : public infini_train::nn::CloneableModule> Forward(const std::vector> &x) override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + private: TransformerConfig config_; int64_t n_head_ = 0; diff --git a/infini_train/include/nn/modules/transformer/mlp.h b/infini_train/include/nn/modules/transformer/mlp.h index ecf5672b6..0248f51ad 100644 --- a/infini_train/include/nn/modules/transformer/mlp.h +++ b/infini_train/include/nn/modules/transformer/mlp.h @@ -19,5 +19,7 @@ class MLP : public infini_train::nn::CloneableModule { std::vector> Forward(const std::vector> &x) override; + + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; }; } // namespace infini_train::nn diff --git a/infini_train/include/nn/modules/transformer/transformer.h b/infini_train/include/nn/modules/transformer/transformer.h index 0471c32fe..40990cb9f 100644 --- a/infini_train/include/nn/modules/transformer/transformer.h +++ b/infini_train/include/nn/modules/transformer/transformer.h @@ -78,6 +78,10 @@ class TransformerModel : public CloneableModule { const TransformerConfig &Config() const { return config_; } + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const 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 d5d527610..45edfd91d 100644 --- a/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h +++ b/infini_train/include/nn/parallel/ddp/distributed_data_parallel.h @@ -31,6 +31,11 @@ class DistributedDataParallel : public nn::Module { std::vector> Forward(const std::vector> &input_tensors) override; std::shared_ptr module() const; + std::vector>> + NamedParameters(const std::string &prefix = "", bool recurse = true, bool remove_duplicate = true) const override; + std::unordered_map> StateDict() const override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + void LoadStateDict(const std::unordered_map> &state_dict) override; std::unique_ptr no_sync() override; diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 25939bdc2..ee717d4ef 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; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + void LoadStateDict(const std::unordered_map> &state_dict) override; + private: void BuildPipelineStage(const std::vector> &recv_shape, Device device, std::vector> &&chunks); diff --git a/infini_train/include/nn/parallel/tensor_parallel.h b/infini_train/include/nn/parallel/tensor_parallel.h index a44aceeaf..d5049877a 100644 --- a/infini_train/include/nn/parallel/tensor_parallel.h +++ b/infini_train/include/nn/parallel/tensor_parallel.h @@ -6,6 +6,7 @@ #include "infini_train/include/autograd/function.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/process_group.h" +#include "infini_train/include/shard_spec.h" namespace infini_train { class Tensor; @@ -37,6 +38,8 @@ class ColumnParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + protected: bool bias_ = true; bool gather_output_ = false; // whether to return full local output tensor after forward (need gather) @@ -66,6 +69,8 @@ class RowParallelLinear : public nn::CloneableModule { bool skip_bias_add() const; bool sequence_parallel() const; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + protected: bool bias_ = true; bool reduce_output_ = false; // whether to return full local output tensor after forward (need reduce) @@ -85,9 +90,13 @@ class VocabParallelEmbedding : public nn::CloneableModule> Forward(const std::vector> &input_tensors) override; + ShardedStateDict BuildShardedStateDict(const std::string &prefix = "") const override; + private: bool reduce_scatter_embeddings_ = false; // whether to perform ReduceScatter after embedding lookup + int64_t vocab_size_global_ = 0; + int64_t embedding_dim_ = 0; int64_t vocab_size_per_partition_ = 0; diff --git a/infini_train/include/optimizer.h b/infini_train/include/optimizer.h index d85b1acea..377924fc3 100644 --- a/infini_train/include/optimizer.h +++ b/infini_train/include/optimizer.h @@ -4,6 +4,7 @@ #include #include #include +#include #include #include #include @@ -52,6 +53,11 @@ class Optimizer { }; namespace optimizers { +inline constexpr std::string_view kAdamOptimizerPrefix = "adam."; +inline constexpr std::string_view kAdamFirstMomentPrefix = "adam.m."; +inline constexpr std::string_view kAdamSecondMomentPrefix = "adam.v."; +inline constexpr std::string_view kAdamStepKey = "adam.t"; + class SGD : public Optimizer { public: SGD(const std::vector> ¶ms, float learning_rate); diff --git a/infini_train/include/shard_spec.h b/infini_train/include/shard_spec.h new file mode 100644 index 000000000..b5af5d24d --- /dev/null +++ b/infini_train/include/shard_spec.h @@ -0,0 +1,68 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/datatype.h" + +namespace infini_train { + +struct ShardSegment { + int64_t global_offset = 0; + int64_t local_offset = 0; + int64_t length = 0; + + bool operator==(const ShardSegment &other) const = default; +}; + +// Logical tensor shard metadata, aligned with Megatron-LM's ShardedTensor model. +struct ShardedTensor { + std::string key; + std::string local_key; + DataType dtype = DataType::kFLOAT32; + std::vector global_shape; + std::vector local_shape; + std::vector global_offset; + std::vector axis_fragmentations; + bool allow_shape_mismatch = false; + // Optional disjoint regions along the single fragmented axis. This is used + // by layouts such as rank-local [Q, K, V], which are not one contiguous + // slice of the logical global [Q, K, V] tensor. + std::vector segments; + + bool operator==(const ShardedTensor &other) const { + return key == other.key && local_key == other.local_key && dtype == other.dtype + && global_shape == other.global_shape && local_shape == other.local_shape + && global_offset == other.global_offset && axis_fragmentations == other.axis_fragmentations + && allow_shape_mismatch == other.allow_shape_mismatch && segments == other.segments; + } +}; + +inline ShardedTensor MakeShardedTensor(std::string key, DataType dtype, std::vector shape) { + const auto ndim = shape.size(); + return {.key = std::move(key), + .dtype = dtype, + .global_shape = shape, + .local_shape = std::move(shape), + .global_offset = std::vector(ndim, 0), + .axis_fragmentations = std::vector(ndim, 1)}; +} + +struct ShardedStateDict { + std::map tensors; + + void Merge(ShardedStateDict &&other) { + for (auto &[key, info] : other.tensors) { + const auto display_key = key; + const auto [_, inserted] = tensors.emplace(std::move(key), std::move(info)); + CHECK(inserted) << "Duplicate sharded state-dict key: " << display_key; + } + } +}; + +} // namespace infini_train diff --git a/infini_train/src/checkpoint/checkpoint.cc b/infini_train/src/checkpoint/checkpoint.cc index ad51bab1c..c02784def 100644 --- a/infini_train/src/checkpoint/checkpoint.cc +++ b/infini_train/src/checkpoint/checkpoint.cc @@ -1,7 +1,11 @@ #include "infini_train/include/checkpoint/checkpoint.h" +#include +#include +#include #include #include +#include #include #include #include @@ -11,8 +15,15 @@ #include "glog/logging.h" +#include "infini_train/include/checkpoint/constants.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/load_strategy.h" +#include "infini_train/include/checkpoint/save_planner.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/parallel_functional.h" +#include "infini_train/include/nn/parallel/work.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -179,75 +190,108 @@ template T ExtractNumberField(const std::string &content, const std } return value; } + +void SynchronizeCheckpointRanks(const nn::Module &model) { + if (nn::parallel::global::GetWorldSize() == 1) { + return; + } + const auto parameters = model.Parameters(); + CHECK(!parameters.empty()) << "Cannot synchronize checkpoint save for a model without parameters"; + auto token = std::make_shared(std::vector{1}, DataType::kFLOAT32, parameters.front()->GetDevice()); + token->Fill(1.0f); + nn::parallel::function::AllReduce(token, nn::parallel::function::ReduceOpType::kSum, nullptr, true)->Synchronize(); +} } // namespace void Checkpoint::Save(const std::filesystem::path &checkpoint_dir, const nn::Module &model, const Optimizer *optimizer, const TrainerState &state, const LRScheduler *lr_scheduler) { std::filesystem::create_directories(checkpoint_dir); - LOG(INFO) << "[CKPT] Save begin: dir=" << checkpoint_dir << ", global_step=" << state.global_step; - - const auto model_path = checkpoint_dir / ("model.ckpt"); - - SaveStateDict(model_path, model.StateDict()); - - if (optimizer != nullptr) { - auto opt_state = optimizer->StateDict(); - if (!opt_state.empty()) { - const auto opt_path = checkpoint_dir / "optimizer.ckpt"; - SaveStateDict(opt_path, opt_state); - } + const int global_rank = nn::parallel::global::thread_global_rank; + const auto staging_root = checkpoint_dir / ".metadata_tmp"; + if (global_rank == 0) { + std::filesystem::remove_all(staging_root); } + SynchronizeCheckpointRanks(model); + + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(global_rank, dp_rank, tp_rank, pp_rank); + // TODO(jym): Select checkpoint writers from each shard's replica_id instead of hard-coding DP rank 0. + if (dp_rank == 0) { + const auto rank_dir = checkpoint_dir / std::format("rank_{:06d}", global_rank); + auto sharded_state = model.BuildShardedStateDict(); + std::unordered_map> optimizer_state; + if (optimizer != nullptr) { + optimizer_state = optimizer->StateDict(); + sharded_state.Merge(checkpoint::BuildOptimizerShardedStateDict(sharded_state, optimizer_state)); + } + const auto write_items = checkpoint::SavePlanner::Plan(sharded_state, global_rank); + SaveLocalShard(rank_dir, sharded_state, write_items, model.StateDict(), optimizer_state, global_rank); - if (lr_scheduler != nullptr) { - SaveLRSchedulerState(checkpoint_dir / "lr_scheduler.ckpt", lr_scheduler->StateDict()); + const auto staging_rank_dir = staging_root / std::format("rank_{:06d}", global_rank); + std::filesystem::create_directories(staging_rank_dir); + std::filesystem::rename(rank_dir / checkpoint::kMetadataFilename, + staging_rank_dir / checkpoint::kMetadataFilename); } - SaveTrainerState(checkpoint_dir / "trainer_state.json", state); - LOG(ERROR) << "[CKPT] Save done: dir=" << checkpoint_dir; + SynchronizeCheckpointRanks(model); + if (global_rank == 0) { + SaveTrainerState(checkpoint_dir / checkpoint::kTrainerStateFilename, state); + if (lr_scheduler != nullptr) { + SaveLRSchedulerState(checkpoint_dir / checkpoint::kLRSchedulerFilename, lr_scheduler->StateDict()); + } + auto metadata = LoadMetadata(staging_root); + CHECK(metadata.has_metadata); + SaveMetadataFile(checkpoint_dir / checkpoint::kTemporaryMetadataFilename, metadata); + if (std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) { + std::filesystem::remove(checkpoint_dir / checkpoint::kMetadataFilename); + } + std::filesystem::rename(checkpoint_dir / checkpoint::kTemporaryMetadataFilename, + checkpoint_dir / checkpoint::kMetadataFilename); + std::filesystem::remove_all(staging_root); + } + SynchronizeCheckpointRanks(model); } void Checkpoint::Load(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, TrainerState &state, LRScheduler *lr_scheduler) { - const auto model_path = checkpoint_dir / "model.ckpt"; - LOG(INFO) << "[CKPT] Loading model: " << model_path; - - model.LoadStateDict(LoadStateDict(model_path)); + const auto metadata = LoadMetadata(checkpoint_dir); + CHECK(metadata.has_metadata); + CHECK_EQ(metadata.version, 3) << "Unsupported distributed checkpoint version: " << metadata.version; + state = LoadTrainerState(checkpoint_dir / checkpoint::kTrainerStateFilename); + + auto model_sharded_state = model.BuildShardedStateDict(); + checkpoint::IndexedRegionLoadStrategy strategy; + auto result = strategy.Execute(checkpoint_dir, checkpoint::LoadPlanner::PlanReshard(metadata, model_sharded_state)); + model.LoadStateDict(result); + + const int current_tp = nn::parallel::global::GetTensorParallelSize(); + const int current_pp = nn::parallel::global::GetPipelineParallelSize(); + state.tp_size = current_tp; + state.pp_size = current_pp; + state.ddp_size = nn::parallel::global::GetDataParallelSize(); + state.sp_size = nn::parallel::global::GetSequenceParallelEnabled() ? current_tp : 1; + state.vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); if (optimizer != nullptr) { - const auto opt_path = checkpoint_dir / "optimizer.ckpt"; - if (std::filesystem::exists(opt_path)) { - LOG(INFO) << "[CKPT] Loading optimizer: " << opt_path; - optimizer->LoadStateDict(LoadStateDict(opt_path)); - } else { - LOG(FATAL) << "Optimizer checkpoint not found at: " << opt_path; - } + auto optimizer_sharded_state + = checkpoint::BuildOptimizerShardedStateDict(model_sharded_state, optimizer->StateDict()); + auto optimizer_plan = checkpoint::LoadPlanner::PlanReshard(metadata, optimizer_sharded_state); + optimizer->LoadStateDict(strategy.Execute(checkpoint_dir, optimizer_plan)); } - - state = LoadTrainerState(checkpoint_dir / "trainer_state.json"); - - if (lr_scheduler != nullptr) { - const auto lr_scheduler_path = checkpoint_dir / "lr_scheduler.ckpt"; - if (std::filesystem::exists(lr_scheduler_path)) { - LOG(INFO) << "[CKPT] Loading LR scheduler: " << lr_scheduler_path; - lr_scheduler->LoadStateDict(LoadLRSchedulerState(lr_scheduler_path)); - } else { - LOG(WARNING) << "[CKPT] LR scheduler checkpoint not found at: " << lr_scheduler_path - << ". Keeping the initialized scheduler state."; - } + if (lr_scheduler != nullptr && std::filesystem::exists(checkpoint_dir / checkpoint::kLRSchedulerFilename)) { + lr_scheduler->LoadStateDict(LoadLRSchedulerState(checkpoint_dir / checkpoint::kLRSchedulerFilename)); } - - LOG(ERROR) << "[CKPT] Load done: global_step=" << state.global_step - << ", consumed_train_samples=" << state.consumed_train_samples << ", topology(ddp,tp,sp,pp)=(" - << state.ddp_size << "," << state.tp_size << "," << state.sp_size << "," << state.pp_size << ")"; } -void Checkpoint::SaveStateDict(const std::filesystem::path &path, - const std::unordered_map> &state_dict) { +Checkpoint::SavedTensorLocations +Checkpoint::SaveStateDict(const std::filesystem::path &path, + const std::unordered_map> &state_dict) { std::ofstream ofs(path, std::ios::binary); CHECK(ofs.is_open()) << "Failed to open checkpoint file: " << path; uint32_t magic = kCkptMagic; uint32_t version = kCkptVersion; + SavedTensorLocations locations; uint32_t count = static_cast(state_dict.size()); ofs.write(reinterpret_cast(&magic), sizeof(magic)); ofs.write(reinterpret_cast(&version), sizeof(version)); @@ -267,8 +311,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) { @@ -318,13 +368,15 @@ void Checkpoint::SaveTrainerState(const std::filesystem::path &path, const Train ofs << " \"n_head\": " << state.n_head << ",\n"; ofs << " \"n_kv_head\": " << state.n_kv_head << ",\n"; ofs << " \"n_embd\": " << state.n_embd << ",\n"; - ofs << " \"vocab_size\": " << state.vocab_size << ",\n"; + ofs << " \"original_vocab_size\": " << state.original_vocab_size << ",\n"; + ofs << " \"padded_vocab_size\": " << state.padded_vocab_size << ",\n"; ofs << " \"global_step\": " << state.global_step << ",\n"; ofs << " \"consumed_train_samples\": " << state.consumed_train_samples << ",\n"; ofs << " \"ddp_size\": " << state.ddp_size << ",\n"; ofs << " \"tp_size\": " << state.tp_size << ",\n"; ofs << " \"sp_size\": " << state.sp_size << ",\n"; - ofs << " \"pp_size\": " << state.pp_size << "\n"; + ofs << " \"pp_size\": " << state.pp_size << ",\n"; + ofs << " \"vpp_size\": " << state.vpp_size << "\n"; ofs << "}\n"; } @@ -339,13 +391,394 @@ TrainerState Checkpoint::LoadTrainerState(const std::filesystem::path &path) { state.n_head = ExtractNumberField(content, "n_head", 0); state.n_kv_head = ExtractNumberField(content, "n_kv_head", 0); state.n_embd = ExtractNumberField(content, "n_embd", 0); - state.vocab_size = ExtractNumberField(content, "vocab_size", 0); + const auto legacy_vocab_size = ExtractNumberField(content, "vocab_size", 0); + state.original_vocab_size = ExtractNumberField(content, "original_vocab_size", 0); + state.padded_vocab_size = ExtractNumberField(content, "padded_vocab_size", legacy_vocab_size); state.global_step = ExtractNumberField(content, "global_step", 0); state.consumed_train_samples = ExtractNumberField(content, "consumed_train_samples", 0); state.ddp_size = ExtractNumberField(content, "ddp_size", 1); state.tp_size = ExtractNumberField(content, "tp_size", 1); state.sp_size = ExtractNumberField(content, "sp_size", 1); state.pp_size = ExtractNumberField(content, "pp_size", 1); + state.vpp_size = ExtractNumberField(content, "vpp_size", 1); return state; } + +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); +} + +static std::string DataTypeToString(DataType dt) { + const auto it = kDataTypeToDesc.find(dt); + CHECK(it != kDataTypeToDesc.end()) << "Unsupported checkpoint tensor dtype: " << static_cast(dt); + return it->second; +} + +void Checkpoint::SaveLocalShard(const std::filesystem::path &checkpoint_dir, const ShardedStateDict &sharded_sd, + const std::vector &write_items, + const std::unordered_map> &state_dict, + const std::unordered_map> &optimizer_state, + int global_rank) { + std::filesystem::create_directories(checkpoint_dir); + LOG(INFO) << "[CKPT] SaveLocalShard begin: dir=" << checkpoint_dir << ", rank=" << global_rank; + + SavedTensorLocations model_file_index; + SavedTensorLocations optimizer_file_index; + + // Save model tensors separately from optimizer tensors. + { + std::unordered_map> filtered_sd; + for (const auto &[key, info] : sharded_sd.tensors) { + // Optimizer tensors are serialized separately. + if (key.starts_with(optimizers::kAdamOptimizerPrefix)) { + continue; + } + // Match metadata keys to the local tensor payloads. + const auto &local_key = info.local_key.empty() ? key : info.local_key; + auto it = state_dict.find(local_key); + if (it != state_dict.end()) { + filtered_sd.emplace(key, it->second); + } + } + if (!filtered_sd.empty()) { + model_file_index = SaveStateDict(checkpoint_dir / checkpoint::kModelCheckpointFilename, filtered_sd); + } + } + + // Save the rank-local optimizer state. + if (!optimizer_state.empty()) { + optimizer_file_index + = SaveStateDict(checkpoint_dir / checkpoint::kOptimizerCheckpointFilename, optimizer_state); + } + + // Write the temporary rank manifest. + { + std::ofstream ofs(checkpoint_dir / checkpoint::kMetadataFilename); + CHECK(ofs.is_open()) << "Failed to open metadata.json: " << checkpoint_dir / checkpoint::kMetadataFilename; + + ofs << "{\n"; + ofs << " \"version\": 3,\n"; + ofs << " \"format\": \"infinitrain_sharded\",\n"; + + ofs << " \"tensors\": [\n"; + + std::vector emitted_items; + for (const auto &item : write_items) { + if (sharded_sd.tensors.contains(item.key)) { + emitted_items.push_back(&item); + } + } + int dp_rank = 0, tp_rank = 0, pp_rank = 0; + nn::parallel::global::GetCoordOf(global_rank, dp_rank, tp_rank, pp_rank); + for (size_t i = 0; i < emitted_items.size(); ++i) { + const auto &item = *emitted_items[i]; + const auto it = sharded_sd.tensors.find(item.key); + const auto &file_index + = item.filename == checkpoint::kOptimizerCheckpointFilename ? optimizer_file_index : model_file_index; + const auto storage_it = file_index.find(item.key); + CHECK(storage_it != file_index.end()) << "Missing stored tensor metadata for " << item.key; + const auto &storage = storage_it->second; + CHECK_EQ(storage.byte_size, item.byte_size); + + ofs << " {\n"; + ofs << " \"key\": \"" << item.key << "\",\n"; + ofs << " \"dtype\": \"" << DataTypeToString(item.dtype) << "\",\n"; + + // global_shape + ofs << " \"global_shape\": ["; + const auto &gs = it->second.global_shape; + for (size_t j = 0; j < gs.size(); ++j) { ofs << gs[j] << (j + 1 < gs.size() ? ", " : ""); } + ofs << "],\n"; + + ofs << " \"local_shape\": ["; + const auto &ls = it->second.local_shape; + for (size_t j = 0; j < ls.size(); ++j) { ofs << ls[j] << (j + 1 < ls.size() ? ", " : ""); } + ofs << "],\n"; + + ofs << " \"global_offset\": ["; + for (size_t j = 0; j < it->second.global_offset.size(); ++j) { + ofs << it->second.global_offset[j] << (j + 1 < it->second.global_offset.size() ? ", " : ""); + } + ofs << "],\n"; + ofs << " \"axis_fragmentations\": ["; + for (size_t j = 0; j < it->second.axis_fragmentations.size(); ++j) { + ofs << it->second.axis_fragmentations[j] << (j + 1 < it->second.axis_fragmentations.size() ? ", " : ""); + } + ofs << "],\n"; + auto write_segments = [&](const char *name, auto member) { + ofs << " \"" << name << "\": ["; + for (size_t j = 0; j < it->second.segments.size(); ++j) { + ofs << it->second.segments[j].*member << (j + 1 < it->second.segments.size() ? ", " : ""); + } + ofs << "],\n"; + }; + write_segments("segment_global_offsets", &ShardSegment::global_offset); + write_segments("segment_local_offsets", &ShardSegment::local_offset); + write_segments("segment_lengths", &ShardSegment::length); + + ofs << " \"file\": \"" << item.filename << "\",\n"; + ofs << " \"offset\": " << storage.data_offset << ",\n"; + ofs << " \"byte_size\": " << item.byte_size << ",\n"; + ofs << " \"pp_rank\": " << pp_rank << ",\n"; + ofs << " \"stored_on_ranks\": [" << global_rank << "]\n"; + ofs << " }"; + if (i + 1 < emitted_items.size()) { + ofs << ","; + } + ofs << "\n"; + } + + ofs << " ]\n"; + ofs << "}\n"; + + LOG(INFO) << "[CKPT] metadata.json written"; + } + + LOG(ERROR) << "[CKPT] SaveLocalShard done: dir=" << checkpoint_dir; +} + +// Load one manifest or aggregate writer manifests while finalizing a checkpoint. +static std::string ExtractJsonString(const std::string &obj, const std::string &key) { + auto token = std::string("\"") + key + "\""; + auto pos = obj.find(token); + if (pos == std::string::npos) { + return ""; + } + auto q1 = obj.find('"', pos + token.size()); + if (q1 == std::string::npos) { + return ""; + } + auto q2 = obj.find('"', q1 + 1); + if (q2 == std::string::npos) { + return ""; + } + return obj.substr(q1 + 1, q2 - q1 - 1); +} + +static std::string_view TrimWhitespace(std::string_view value) { + while (!value.empty() && std::isspace(static_cast(value.front()))) { value.remove_prefix(1); } + while (!value.empty() && std::isspace(static_cast(value.back()))) { value.remove_suffix(1); } + return value; +} + +template std::vector ExtractIntegerArray(const std::string &obj, const char *field_name) { + const auto field = std::string("\"") + field_name + "\""; + const auto field_pos = obj.find(field); + if (field_pos == std::string::npos) { + return {}; + } + + const auto array_start = obj.find('[', field_pos + field.size()); + const auto array_end = obj.find(']', array_start); + CHECK(array_start != std::string::npos && array_end != std::string::npos) + << "Invalid integer array in checkpoint metadata field " << field_name; + + std::string_view contents(obj.data() + array_start + 1, array_end - array_start - 1); + contents = TrimWhitespace(contents); + if (contents.empty()) { + return {}; + } + + std::vector values; + while (!contents.empty()) { + const auto delimiter = contents.find(','); + const auto token = TrimWhitespace(contents.substr(0, delimiter)); + CHECK(!token.empty()) << "Empty integer in checkpoint metadata field " << field_name; + + T value{}; + const auto [end, error] = std::from_chars(token.data(), token.data() + token.size(), value); + CHECK(error == std::errc{} && end == token.data() + token.size()) + << "Invalid integer in checkpoint metadata field " << field_name << ": " << token; + values.push_back(value); + + if (delimiter == std::string_view::npos) { + break; + } + contents.remove_prefix(delimiter + 1); + contents = TrimWhitespace(contents); + CHECK(!contents.empty()) << "Trailing comma in checkpoint metadata field " << field_name; + } + return values; +} + +static Checkpoint::CheckpointMetadata LoadSingleMetadata(const std::filesystem::path &checkpoint_dir) { + Checkpoint::CheckpointMetadata meta; + auto metadata_path = checkpoint_dir / checkpoint::kMetadataFilename; + 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); + + // Locate the tensors array. + auto tensors_key = content.find("\"tensors\""); + if (tensors_key == std::string::npos) { + return meta; + } + + auto array_start = content.find('[', tensors_key); + if (array_start == std::string::npos) { + return meta; + } + + int depth = 1; + size_t pos = array_start + 1; + while (pos < content.size() && depth > 0) { + if (content[pos] == '[') { + ++depth; + } else if (content[pos] == ']') { + --depth; + } + ++pos; + } + std::string tensor_block = content.substr(array_start + 1, pos - array_start - 2); + + // Parse each tensor object. + size_t obj_pos = 0; + while ((obj_pos = tensor_block.find('{', obj_pos)) != std::string::npos) { + int object_depth = 1; + size_t obj_end = obj_pos + 1; + while (obj_end < tensor_block.size() && object_depth > 0) { + if (tensor_block[obj_end] == '{') { + ++object_depth; + } + if (tensor_block[obj_end] == '}') { + --object_depth; + } + ++obj_end; + } + if (obj_end > 0) { + --obj_end; + } + if (obj_end == std::string::npos) { + break; + } + + std::string obj = tensor_block.substr(obj_pos, obj_end - obj_pos + 1); + + Checkpoint::CheckpointMetadata::TensorEntry entry; + entry.key = ExtractJsonString(obj, "key"); + entry.file = ExtractJsonString(obj, "file"); + entry.dtype_str = ExtractJsonString(obj, "dtype"); + entry.offset = ExtractNumberField(obj, "offset", 0); + entry.byte_size = ExtractNumberField(obj, "byte_size", 0); + entry.pp_rank = ExtractNumberField(obj, "pp_rank", 0); + + entry.global_shape = ExtractIntegerArray(obj, "global_shape"); + entry.local_shape = ExtractIntegerArray(obj, "local_shape"); + entry.global_offset = ExtractIntegerArray(obj, "global_offset"); + entry.axis_fragmentations = ExtractIntegerArray(obj, "axis_fragmentations"); + + const auto segment_global_offsets = ExtractIntegerArray(obj, "segment_global_offsets"); + const auto segment_local_offsets = ExtractIntegerArray(obj, "segment_local_offsets"); + const auto segment_lengths = ExtractIntegerArray(obj, "segment_lengths"); + CHECK_EQ(segment_global_offsets.size(), segment_local_offsets.size()); + CHECK_EQ(segment_global_offsets.size(), segment_lengths.size()); + for (size_t i = 0; i < segment_lengths.size(); ++i) { + entry.segments.push_back({.global_offset = segment_global_offsets[i], + .local_offset = segment_local_offsets[i], + .length = segment_lengths[i]}); + } + + entry.stored_on_ranks = ExtractIntegerArray(obj, "stored_on_ranks"); + + meta.tensors.push_back(std::move(entry)); + obj_pos = obj_end + 1; + } + + LOG(INFO) << "[CKPT] Loaded metadata.json: " << meta.tensors.size() << " tensor shards"; + return meta; +} + +Checkpoint::CheckpointMetadata Checkpoint::LoadMetadata(const std::filesystem::path &checkpoint_dir) { + if (std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) { + return LoadSingleMetadata(checkpoint_dir); + } + + CheckpointMetadata merged; + for (const auto &entry : std::filesystem::directory_iterator(checkpoint_dir)) { + if (!entry.is_directory() || !entry.path().filename().string().starts_with("rank_") + || !std::filesystem::exists(entry.path() / checkpoint::kMetadataFilename)) { + 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 << " \"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", &ShardSegment::global_offset); + write_segments("segment_local_offsets", &ShardSegment::local_offset); + write_segments("segment_lengths", &ShardSegment::length); + ofs << " \"file\": \"" << tensor.file << "\",\n"; + ofs << " \"offset\": " << tensor.offset << ",\n"; + ofs << " \"byte_size\": " << tensor.byte_size << ",\n"; + ofs << " \"pp_rank\": " << tensor.pp_rank << ",\n"; + ofs << " \"stored_on_ranks\": ["; + for (size_t r = 0; r < tensor.stored_on_ranks.size(); ++r) { + ofs << tensor.stored_on_ranks[r] << (r + 1 < tensor.stored_on_ranks.size() ? ", " : ""); + } + ofs << "]\n"; + ofs << " }" << (i + 1 < metadata.tensors.size() ? "," : "") << "\n"; + } + ofs << " ]\n"; + ofs << "}\n"; + CHECK(ofs.good()) << "Failed while writing checkpoint metadata: " << path; +} } // namespace infini_train diff --git a/infini_train/src/checkpoint/checkpoint_manager.cc b/infini_train/src/checkpoint/checkpoint_manager.cc index c6e31cddb..512eacd9f 100644 --- a/infini_train/src/checkpoint/checkpoint_manager.cc +++ b/infini_train/src/checkpoint/checkpoint_manager.cc @@ -1,48 +1,55 @@ #include "infini_train/include/checkpoint/checkpoint_manager.h" -#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/constants.h" +#include "infini_train/include/lr_scheduler.h" +#include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" -#include "infini_train/include/nn/parallel/global.h" -#include "infini_train/include/tensor.h" +#include "infini_train/include/nn/parallel/ddp/distributed_optimizer.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. +namespace { + +std::filesystem::path ResolveCheckpointDirectory(const std::filesystem::path &root) { + const auto latest_path = root / checkpoint::kLatestIterationFilename; + if (!std::filesystem::exists(latest_path)) { + return root; + } + std::ifstream latest(latest_path); + int64_t iteration = 0; + latest >> iteration; + const auto directory = root / std::format("iter_{:07d}", iteration); + CHECK(std::filesystem::exists(directory)) << "Latest checkpoint directory does not exist: " << directory; + return directory; +} + +} // namespace + ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs &args) { ResumeFromCheckpointResult result; if (args.resume_root.empty()) { LOG(INFO) << "No checkpoint specified for resume. Starting training from scratch."; return result; } + CHECK(dynamic_cast(args.optimizer.get()) == nullptr) + << "Checkpoint restore does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; - 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 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; - } - } - - Checkpoint::Load(resume_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get()); - - result.global_step = static_cast(args.state.global_step); + const auto checkpoint_dir = ResolveCheckpointDirectory(args.resume_root); + CHECK(std::filesystem::exists(checkpoint_dir / checkpoint::kMetadataFilename)) + << "Checkpoint metadata.json not found: " << checkpoint_dir; + Checkpoint::Load(checkpoint_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get()); CHECK_EQ(args.state.n_layer, args.model_config.n_layer) << "n_layer mismatch: ckpt=" << args.state.n_layer << ", config=" << args.model_config.n_layer; @@ -52,72 +59,85 @@ ResumeFromCheckpointResult ResumeFromCheckpoint(const ResumeFromCheckpointArgs & << "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; - + if (args.state.original_vocab_size > 0) { + CHECK_EQ(args.state.original_vocab_size, args.model_config.original_vocab_size) + << "original_vocab_size mismatch: ckpt=" << args.state.original_vocab_size + << ", config=" << args.model_config.original_vocab_size; + CHECK_GE(args.state.padded_vocab_size, args.state.original_vocab_size) + << "Checkpoint padded vocabulary cannot represent its logical vocabulary"; + } else { + // Legacy trainer_state.json only stored the padded vocab size. + CHECK_GE(args.state.padded_vocab_size, args.model_config.original_vocab_size) + << "Legacy checkpoint vocabulary cannot represent the configured logical vocabulary"; + } + result.global_step = static_cast(args.state.global_step); result.consumed_train_samples = static_cast(std::max(args.state.consumed_train_samples, 0)); if (args.rank.IsMainRank()) { LOG(INFO) << std::format("Resume training from step {}, consumed_train_samples {}", args.state.global_step, args.state.consumed_train_samples); } - return result; } void SaveCheckpoint(const SaveCheckpointArgs &args) { - const auto ckpt_start = std::chrono::high_resolution_clock::now(); - - TrainerState state; - state.global_step = args.global_step; - state.consumed_train_samples = static_cast(args.consumed_train_samples); - state.n_layer = args.n_layer; - state.n_head = args.n_head; - state.n_kv_head = args.n_kv_head; - state.n_embd = args.n_embd; - state.vocab_size = args.vocab_size; - state.ddp_size = args.ddp_size; - state.tp_size = args.tp_size; - state.sp_size = args.sp_size; - state.pp_size = args.pp_size; - - Checkpoint::Save(args.save_dir, args.model, args.optimizer, state, args.lr_scheduler); - - const auto ckpt_end = std::chrono::high_resolution_clock::now(); - const double ckpt_ms = std::chrono::duration(ckpt_end - ckpt_start).count(); - - if (!args.rank.IsMainRank()) { - return; + CHECK(dynamic_cast(args.optimizer) == nullptr) + << "Checkpoint save does not support DistributedOptimizer/ZeRO optimizer state; use zero_stage=0"; + CHECK_GT(args.original_vocab_size, 0); + CHECK_GE(args.padded_vocab_size, args.original_vocab_size); + const auto checkpoint_start = std::chrono::high_resolution_clock::now(); + // Snapshot training progress and the topology that produced this checkpoint. + TrainerState state{.global_step = args.global_step, + .consumed_train_samples = static_cast(args.consumed_train_samples), + .n_layer = args.n_layer, + .n_head = args.n_head, + .n_kv_head = args.n_kv_head, + .n_embd = args.n_embd, + .original_vocab_size = args.original_vocab_size, + .padded_vocab_size = args.padded_vocab_size, + .ddp_size = args.ddp_size, + .tp_size = args.tp_size, + .sp_size = args.sp_size, + .pp_size = args.pp_size, + .vpp_size = args.vpp_size}; + const auto iteration_dir = args.checkpoint_root_dir.empty() + ? args.save_dir + : args.checkpoint_root_dir / std::format("iter_{:07d}", args.global_step); + Checkpoint::Save(iteration_dir, args.model, args.optimizer, state, args.lr_scheduler); + + if (args.rank.IsMainRank() && !args.checkpoint_root_dir.empty()) { + const auto latest = args.checkpoint_root_dir / checkpoint::kLatestIterationFilename; + const auto temporary_latest = args.checkpoint_root_dir / checkpoint::kTemporaryLatestIterationFilename; + { + 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); } - 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. - 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("checkpoint_step_")) { - ckpts.push_back(entry.path()); + if (entry.is_directory() && entry.path().filename().string().starts_with("iter_")) { + checkpoints.push_back(entry.path()); } } - std::sort(ckpts.begin(), ckpts.end()); - while (ckpts.size() > args.max_checkpoint_keep) { - std::filesystem::remove_all(ckpts.front()); - ckpts.erase(ckpts.begin()); + // FIXME(jym): Pruning relies on lexicographic sorting of checkpoint directory names. + // This is only correct while iteration directories use zero-padded names (e.g. iter_0000042). + // If the naming convention changes to unpadded names, parse the iteration and sort numerically instead. + std::sort(checkpoints.begin(), checkpoints.end()); + while (checkpoints.size() > args.max_checkpoint_keep) { + std::filesystem::remove_all(checkpoints.front()); + checkpoints.erase(checkpoints.begin()); } } + + const auto checkpoint_end = std::chrono::high_resolution_clock::now(); + const double elapsed_ms = std::chrono::duration(checkpoint_end - checkpoint_start).count(); + LOG(INFO) << std::format("Checkpoint saved at: {} ({:.2f} ms)", iteration_dir.string(), elapsed_ms); } size_t DataLoaderBatchesToSkip(size_t consumed_train_samples, size_t local_batch_size, size_t ddp_world_size) { diff --git a/infini_train/src/checkpoint/load_planner.cc b/infini_train/src/checkpoint/load_planner.cc new file mode 100644 index 000000000..708171955 --- /dev/null +++ b/infini_train/src/checkpoint/load_planner.cc @@ -0,0 +1,241 @@ +#include "infini_train/include/checkpoint/load_planner.h" + +#include +#include +#include +#include + +#include "glog/logging.h" + +namespace infini_train::checkpoint { +namespace { + +DataType StringToDataType(const std::string &value) { + static const std::unordered_map legacy_names = { + {"bfloat16", DataType::kBFLOAT16}, + {"float16", DataType::kFLOAT16}, + {"float32", DataType::kFLOAT32}, + {"float64", DataType::kFLOAT64}, + }; + if (const auto it = legacy_names.find(value); it != legacy_names.end()) { + return it->second; + } + for (const auto &[dtype, description] : kDataTypeToDesc) { + if (description == value) { + return dtype; + } + } + LOG(FATAL) << "Unsupported checkpoint tensor dtype: " << value; + return DataType::kFLOAT32; +} + +int FragmentedAxis(const std::vector &axis_fragmentations) { + int fragmented_axis = -1; + for (size_t dim = 0; dim < axis_fragmentations.size(); ++dim) { + if (axis_fragmentations[dim] <= 1) { + continue; + } + CHECK_EQ(fragmented_axis, -1) << "Multi-axis checkpoint sharding is not supported yet"; + fragmented_axis = static_cast(dim); + } + return fragmented_axis; +} + +int ResolveShardDim(const std::string &key, const ShardedTensor &target, int saved_axis, + const std::vector &saved_global_shape) { + const int target_axis = FragmentedAxis(target.axis_fragmentations); + int shard_dim = target_axis >= 0 ? target_axis : saved_axis; + if (shard_dim < 0 && saved_global_shape != target.global_shape) { + shard_dim = 0; + } + if (saved_axis >= 0 && target_axis >= 0) { + CHECK_EQ(saved_axis, target_axis) << "Shard dimension changed for tensor " << key; + } + return shard_dim; +} + +bool IsPaddingCompatible(bool allow_shape_mismatch, const std::vector &source, + const std::vector &target) { + if (!allow_shape_mismatch || source.size() != target.size() || source.empty()) { + return false; + } + // Dim 0 is the padded vocabulary axis and may grow or shrink across TP layouts. + // Every non-padding dimension must still describe the same tensor shape. + for (size_t dim = 1; dim < source.size(); ++dim) { + if (source[dim] != target[dim]) { + return false; + } + } + return true; +} + +void ValidateCoordinates(const std::string &key, const std::vector &global_shape, + const std::vector &local_shape, const std::vector &global_offset, + const std::vector &axis_fragmentations) { + CHECK_EQ(global_shape.size(), local_shape.size()) << "Invalid local rank for tensor " << key; + CHECK_EQ(global_shape.size(), global_offset.size()) << "Invalid offset rank for tensor " << key; + CHECK_EQ(global_shape.size(), axis_fragmentations.size()) << "Invalid fragmentation rank for tensor " << key; + for (size_t dim = 0; dim < global_shape.size(); ++dim) { + CHECK_GT(global_shape[dim], 0) << "Invalid global shape for tensor " << key; + CHECK_GT(local_shape[dim], 0) << "Invalid local shape for tensor " << key; + CHECK_GE(global_offset[dim], 0) << "Invalid global offset for tensor " << key; + CHECK_LE(global_offset[dim] + local_shape[dim], global_shape[dim]) + << "Shard exceeds global shape for tensor " << key; + CHECK_GE(axis_fragmentations[dim], 1) << "Invalid axis fragmentation for tensor " << key; + } +} + +} // namespace + +LoadPlan LoadPlanner::PlanReshard(const Checkpoint::CheckpointMetadata &metadata, + const ShardedStateDict &target_state_dict) { + LoadPlan plan; + for (const auto &[key, target] : target_state_dict.tensors) { + ValidateCoordinates(key, target.global_shape, target.local_shape, target.global_offset, + target.axis_fragmentations); + TargetTensorPlan tensor_plan{.key = key, + .dtype = target.dtype, + .global_shape = target.global_shape, + .target_shape = target.local_shape, + .shard_dim = FragmentedAxis(target.axis_fragmentations)}; + + std::vector candidates; + for (const auto &entry : metadata.tensors) { + if (entry.key == key) { + candidates.push_back(&entry); + } + } + CHECK(!candidates.empty()) << "No saved shard found for target tensor: " << key; + const int saved_axis = FragmentedAxis(candidates.front()->axis_fragmentations); + for (const auto *source : candidates) { + ValidateCoordinates(key, source->global_shape, source->local_shape, source->global_offset, + source->axis_fragmentations); + CHECK(source->global_shape == target.global_shape + || IsPaddingCompatible(target.allow_shape_mismatch, source->global_shape, target.global_shape)) + << "Global shape changed for tensor " << key; + CHECK_EQ(FragmentedAxis(source->axis_fragmentations), saved_axis) + << "Inconsistent saved shard dimensions for tensor " << key; + } + + 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; + } + + tensor_plan.shard_dim = ResolveShardDim(key, target, saved_axis, candidates.front()->global_shape); + + if (tensor_plan.shard_dim < 0) { + const auto *source = candidates.front(); + tensor_plan.reads.push_back({.key = key, + .filename = source->file, + .dtype = StringToDataType(source->dtype_str), + .global_shape = source->global_shape, + .byte_size = source->byte_size, + .data_offset = source->offset, + .shard_dim = -1, + .source_shape = source->local_shape}); + plan.tensors.emplace(key, std::move(tensor_plan)); + continue; + } + + const int dim = tensor_plan.shard_dim; + const int64_t target_start = target.global_offset[dim]; + const int64_t target_length = target.local_shape[dim]; + const int64_t target_end = target_start + target_length; + std::set> seen_source_ranges; + + for (const auto *source : candidates) { + const int64_t saved_start = source->global_offset[dim]; + const int64_t saved_length = source->local_shape[dim]; + if (!seen_source_ranges.emplace(saved_start, saved_length).second) { + continue; + } + const int64_t overlap_start = std::max(saved_start, target_start); + const int64_t overlap_end = std::min(saved_start + saved_length, target_end); + if (overlap_start >= overlap_end) { + continue; + } + + tensor_plan.reads.push_back({.key = key, + .filename = source->file, + .dtype = StringToDataType(source->dtype_str), + .global_shape = source->global_shape, + .byte_size = source->byte_size, + .data_offset = source->offset, + .shard_dim = dim, + .source_offset = overlap_start - saved_start, + .target_offset = overlap_start - target_start, + .length = overlap_end - overlap_start, + .source_shape = source->local_shape}); + } + + std::sort(tensor_plan.reads.begin(), tensor_plan.reads.end(), + [](const ReadItem &left, const ReadItem &right) { return left.target_offset < right.target_offset; }); + int64_t covered = 0; + for (const auto &read : tensor_plan.reads) { + CHECK_EQ(read.target_offset, covered) << "Gap or overlap in target shard plan for " << key; + covered += read.length; + } + if (covered < target_length) { + CHECK(target.allow_shape_mismatch) << "Incomplete target shard plan for " << key; + CHECK_EQ(dim, 0) << "Vocabulary padding is only supported along dim 0"; + CHECK_EQ(target_start + covered, candidates.front()->global_shape[0]) + << "Only trailing vocabulary padding is supported for " << key; + tensor_plan.trailing_zero_fill = target_length - covered; + covered = target_length; + } + CHECK_EQ(covered, target_length) << "Incomplete target shard plan for " << key; + plan.tensors.emplace(key, std::move(tensor_plan)); + } + return plan; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/checkpoint/load_strategy.cc b/infini_train/src/checkpoint/load_strategy.cc new file mode 100644 index 000000000..35e967c3b --- /dev/null +++ b/infini_train/src/checkpoint/load_strategy.cc @@ -0,0 +1,133 @@ +#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); + auto piece = read.shard_dim < 0 ? ReadTensor(stream, read) : ReadTensorRegion(stream, read); + if (piece->Dtype() != tensor_plan.dtype) { + piece = std::make_shared(piece->To(tensor_plan.dtype)); + } + pieces.push_back(std::move(piece)); + } + + if (tensor_plan.trailing_zero_fill > 0) { + 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/save_planner.cc b/infini_train/src/checkpoint/save_planner.cc new file mode 100644 index 000000000..3a8004b28 --- /dev/null +++ b/infini_train/src/checkpoint/save_planner.cc @@ -0,0 +1,65 @@ +#include "infini_train/include/checkpoint/save_planner.h" + +#include "glog/logging.h" + +#include "infini_train/include/checkpoint/constants.h" +#include "infini_train/include/optimizer.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::checkpoint { + +ShardedStateDict +BuildOptimizerShardedStateDict(const ShardedStateDict &model_state, + const std::unordered_map> &optimizer_state) { + ShardedStateDict result; + for (const auto &[key, tensor] : optimizer_state) { + if (key == optimizers::kAdamStepKey) { + auto info = MakeShardedTensor(key, tensor->Dtype(), tensor->Dims()); + info.local_key = key; + result.tensors.emplace(key, std::move(info)); + continue; + } + + std::string parameter_key; + if (key.starts_with(optimizers::kAdamFirstMomentPrefix)) { + parameter_key = key.substr(optimizers::kAdamFirstMomentPrefix.size()); + } else if (key.starts_with(optimizers::kAdamSecondMomentPrefix)) { + parameter_key = key.substr(optimizers::kAdamSecondMomentPrefix.size()); + } else { + CHECK(false) << "Unsupported optimizer state key: " << key; + } + + 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 named parameters."; + auto info = model_it->second; + info.key = key; + info.local_key = key; + info.dtype = tensor->Dtype(); + result.tensors.emplace(key, std::move(info)); + } + return result; +} + +std::vector SavePlanner::Plan(const ShardedStateDict &sd, int rank) { + std::vector items; + for (auto &[key, info] : sd.tensors) { + bool is_optimizer = key.starts_with(optimizers::kAdamOptimizerPrefix); + WriteItem item; + item.key = key; + item.filename = is_optimizer ? kOptimizerCheckpointFilename : kModelCheckpointFilename; + item.byte_size = TensorByteSize(info.dtype, info.local_shape); + item.dtype = info.dtype; + item.local_shape = info.local_shape; + item.global_offset = info.global_offset; + item.axis_fragmentations = info.axis_fragmentations; + item.rank = rank; + + items.push_back(std::move(item)); + } + + return items; +} + +} // namespace infini_train::checkpoint diff --git a/infini_train/src/nn/lora/lora_parallel_linear.cc b/infini_train/src/nn/lora/lora_parallel_linear.cc index 9b038e2d4..45130faf4 100644 --- a/infini_train/src/nn/lora/lora_parallel_linear.cc +++ b/infini_train/src/nn/lora/lora_parallel_linear.cc @@ -219,6 +219,27 @@ std::vector> LoRAColumnParallelLinear::LoRAParameters() return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } +ShardedStateDict LoRAColumnParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + auto state = parallel::ColumnParallelLinear::BuildShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + + const auto &lora_a = parameter(kParamLoraAName); + const auto a_key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; + auto a = MakeShardedTensor(a_key, lora_a->Dtype(), lora_a->Dims()); + state.tensors.emplace(a.key, std::move(a)); + + const auto &lora_b = parameter(kParamLoraBName); + ShardedTensor b; + b.key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; + b.dtype = lora_b->Dtype(); + b.global_shape = {lora_b->Dims()[0] * tp_size, lora_b->Dims()[1]}; + b.local_shape = lora_b->Dims(); + b.global_offset = {lora_b->Dims()[0] * parallel::tp_rank, 0}; + b.axis_fragmentations = {tp_size, 1}; + state.tensors.emplace(b.key, std::move(b)); + return state; +} + bool LoRAColumnParallelLinear::IsMerged() const { return merged_; } int64_t LoRAColumnParallelLinear::in_features() const { return in_features_; } @@ -429,6 +450,27 @@ std::vector> LoRARowParallelLinear::LoRAParameters() con return {parameters_.at(kParamLoraAName), parameters_.at(kParamLoraBName)}; } +ShardedStateDict LoRARowParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + auto state = parallel::RowParallelLinear::BuildShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + + const auto &lora_a = parameter(kParamLoraAName); + ShardedTensor a; + a.key = prefix.empty() ? kParamLoraAName : prefix + "." + kParamLoraAName; + a.dtype = lora_a->Dtype(); + a.global_shape = {lora_a->Dims()[0], lora_a->Dims()[1] * tp_size}; + a.local_shape = lora_a->Dims(); + a.global_offset = {0, lora_a->Dims()[1] * parallel::tp_rank}; + a.axis_fragmentations = {1, tp_size}; + state.tensors.emplace(a.key, std::move(a)); + + const auto &lora_b = parameter(kParamLoraBName); + const auto b_key = prefix.empty() ? kParamLoraBName : prefix + "." + kParamLoraBName; + auto b = MakeShardedTensor(b_key, lora_b->Dtype(), lora_b->Dims()); + state.tensors.emplace(b.key, std::move(b)); + return state; +} + bool LoRARowParallelLinear::IsMerged() const { return merged_; } int64_t LoRARowParallelLinear::in_features() const { return in_features_; } diff --git a/infini_train/src/nn/modules/module.cc b/infini_train/src/nn/modules/module.cc index 498bcc075..7ad6f6624 100644 --- a/infini_train/src/nn/modules/module.cc +++ b/infini_train/src/nn/modules/module.cc @@ -188,6 +188,32 @@ std::unordered_map> Module::StateDict() con return state; } +ShardedStateDict Module::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; + + for (auto &[name, param] : parameters_) { + auto key = prefix.empty() ? name : prefix + "." + name; + sd.tensors.emplace(key, MakeShardedTensor(key, param->Dtype(), param->Dims())); + } + + for (auto &[name, buffer] : buffers_) { + auto key = prefix.empty() ? name : prefix + "." + name; + sd.tensors.emplace(key, MakeShardedTensor(key, buffer->Dtype(), buffer->Dims())); + } + + for (auto &[name, module] : modules_) { + if (name.starts_with("__pp")) { + continue; + } + + auto child_prefix = prefix.empty() ? name : prefix + "." + name; + auto child_sd = module->BuildShardedStateDict(child_prefix); + sd.Merge(std::move(child_sd)); + } + + return sd; +} + void Module::LoadStateDict(const std::unordered_map> &state_dict) { // Stage 1: Validate all keys, shapes, and dtypes without copying std::vector error_msgs; diff --git a/infini_train/src/nn/modules/transformer/causal_self_attention.cc b/infini_train/src/nn/modules/transformer/causal_self_attention.cc index 31676b1b4..0d2f48759 100644 --- a/infini_train/src/nn/modules/transformer/causal_self_attention.cc +++ b/infini_train/src/nn/modules/transformer/causal_self_attention.cc @@ -10,6 +10,7 @@ #include "infini_train/include/nn/functional.h" #include "infini_train/include/nn/init.h" +#include "infini_train/include/nn/lora/lora_parallel_linear.h" #include "infini_train/include/nn/modules/normalization.h" #include "infini_train/include/nn/modules/sparse.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" @@ -85,6 +86,43 @@ void CausalSelfAttention::SetupAttention(const TransformerConfig &config) { } } +ShardedStateDict CausalSelfAttention::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); + const int tp_size = parallel::global::GetTensorParallelSize(); + const int rank = parallel::tp_rank; + const int64_t q_global = n_head_ * head_dim_; + 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; + + // FIXME(jym): Transformer should not depend on LoRA just to identify lora_B. Move the packed-QKV segment layout + // into a dedicated sharding abstraction shared by the base attention weight and LoRA parameters. + const auto c_attn_prefix = prefix.empty() ? kCAttnLayerName : prefix + "." + kCAttnLayerName; + auto set_qkv_segments = [&](const std::string ¶meter_name) { + const auto key = c_attn_prefix + "." + parameter_name; + 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); + } + set_qkv_segments(lora::LoRAColumnParallelLinear::kParamLoraBName); + 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/mlp.cc b/infini_train/src/nn/modules/transformer/mlp.cc index 3cf94f91c..8c12c0cfa 100644 --- a/infini_train/src/nn/modules/transformer/mlp.cc +++ b/infini_train/src/nn/modules/transformer/mlp.cc @@ -93,4 +93,39 @@ MLP::Forward(const std::vector> &x) { return (*modules_[kCProjLayerName])(x2); } +ShardedStateDict MLP::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); + if (!modules_.contains(kSwiGLULayerName)) { + return state; + } + + const auto c_fc_prefix = prefix.empty() ? kCFcLayerName : prefix + "." + kCFcLayerName; + const auto weight_key = c_fc_prefix + "." + parallel::ColumnParallelLinear::kParamWeightName; + const auto weight_it = state.tensors.find(weight_key); + CHECK(weight_it != state.tensors.end()) << "Missing packed SwiGLU weight metadata: " << weight_key; + + const auto &weight = weight_it->second; + CHECK(!weight.global_shape.empty() && !weight.local_shape.empty()); + CHECK_EQ(weight.global_shape[0] % 2, 0); + CHECK_EQ(weight.local_shape[0] % 2, 0); + const int64_t global_rows_per_projection = weight.global_shape[0] / 2; + const int64_t local_rows_per_projection = weight.local_shape[0] / 2; + const int rank = parallel::tp_rank; + + for (auto &[key, tensor] : state.tensors) { + if (!key.starts_with(c_fc_prefix + ".") || tensor.global_shape.empty() || tensor.local_shape.empty() + || tensor.global_shape[0] != weight.global_shape[0] || tensor.local_shape[0] != weight.local_shape[0]) { + continue; + } + tensor.global_offset.assign(tensor.global_shape.size(), 0); + tensor.segments = { + {.global_offset = rank * local_rows_per_projection, .local_offset = 0, .length = local_rows_per_projection}, + {.global_offset = global_rows_per_projection + rank * local_rows_per_projection, + .local_offset = local_rows_per_projection, + .length = local_rows_per_projection}, + }; + } + return state; +} + } // namespace infini_train::nn diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index f1048058a..d39a13483 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -237,17 +237,19 @@ TransformerModel::TransformerModel(const TransformerConfig config) auto chunk = std::make_shared(config_, start_layer, end_layer); start_layer_to_layer_size_and_chunk[start_layer] = std::make_pair(end_layer - start_layer, chunk); } - std::vector> h; + std::unordered_map> h; int chunk_idx = 0; for (auto &[start_layer, layer_size_and_chunk] : start_layer_to_layer_size_and_chunk) { auto [layer_size, chunk] = layer_size_and_chunk; - for (int idx = 0; idx < layer_size; ++idx) { - h.push_back(chunk->mutable_module(TransformerChunk::kHLayerName)->mutable_module(std::to_string(idx))); + for (int local_layer = 0; local_layer < layer_size; ++local_layer) { + const auto global_layer = start_layer + local_layer; + h[std::to_string(global_layer)] + = chunk->mutable_module(TransformerChunk::kHLayerName)->mutable_module(std::to_string(local_layer)); } modules_[kPPChunkNamePrefix + std::to_string(chunk_idx)] = std::move(chunk); ++chunk_idx; } - transformer[TransformerChunk::kHLayerName] = std::make_shared(std::move(h)); + transformer[TransformerChunk::kHLayerName] = std::make_shared(std::move(h)); } if (stage_info_.is_last_stage) { @@ -265,6 +267,10 @@ TransformerModel::TransformerModel(const TransformerConfig config) // applied after loading weights so it won't be overwritten. Also fix GPT2::FromLLMC() loading logic to respect // weight tying (do not create/load a separate lm_head.weight tensor; load once into the tied weight) so // parameter counting matches PyTorch/PEFT. + // FIXME(jym): A checkpoint saved with PP > 1 contains independent wte.weight and lm_head.weight tensors. Loading + // it with PP == 1 aliases both keys to the same destination Tensor, so Module::LoadStateDict copies both values + // and the final result depends on unordered-map iteration order. Canonicalize tied checkpoint keys, or reject + // PP > 1 to PP == 1 resharding for tied models until cross-stage weight tying is supported. if (config_.tie_weights && nn::parallel::global::GetPipelineParallelSize() == 1) { // https://paperswithcode.com/method/weight-tying *mutable_module(kTransformerModelName) @@ -275,6 +281,43 @@ TransformerModel::TransformerModel(const TransformerConfig config) } } +ShardedStateDict TransformerModel::BuildShardedStateDict(const std::string &prefix) const { + auto state = Module::BuildShardedStateDict(prefix); + if (stage_info_.is_last_stage) { + const auto lm_head_prefix = prefix.empty() ? TransformerLastStage::kLMHeadLayerName + : prefix + "." + TransformerLastStage::kLMHeadLayerName; + const auto weight_key = lm_head_prefix + "." + parallel::ColumnParallelLinear::kParamWeightName; + state.tensors.at(weight_key).allow_shape_mismatch = true; + } + return state; +} + +std::vector>> +TransformerModel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + if (!recurse) { + return Module::NamedParameters(prefix, false, remove_duplicate); + } + + // Select public aliases so optimizer state keys match ShardedStateDict keys. + auto parameters = Module::NamedParameters(prefix, true, false); + const auto sharded_state = BuildShardedStateDict(prefix); + + std::vector>> result; + std::unordered_set visited; + const auto private_pipeline_prefix = prefix.empty() ? "__pp" : prefix + ".__pp"; + for (auto &[name, parameter] : parameters) { + if (name.starts_with(private_pipeline_prefix)) { + continue; + } + CHECK(sharded_state.tensors.contains(name)) << "Parameter is missing from sharded state dict: " << name; + if (remove_duplicate && !visited.insert(parameter.get()).second) { + continue; + } + result.emplace_back(std::move(name), std::move(parameter)); + } + return result; +} + 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 e063fcb6f..8878944a4 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -211,6 +211,24 @@ void DistributedDataParallel::OnGradReady(const std::shared_ptr ¶m) } } +std::vector>> +DistributedDataParallel::NamedParameters(const std::string &prefix, bool recurse, bool remove_duplicate) const { + return modules_.at(kModuleName)->NamedParameters(prefix, recurse, remove_duplicate); +} + +std::unordered_map> DistributedDataParallel::StateDict() const { + return modules_.at(kModuleName)->StateDict(); +} + +ShardedStateDict DistributedDataParallel::BuildShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->BuildShardedStateDict(prefix); +} + +void DistributedDataParallel::LoadStateDict( + const std::unordered_map> &state_dict) { + modules_.at(kModuleName)->LoadStateDict(state_dict); +} + std::vector> DistributedDataParallel::Forward(const std::vector> &input_tensors) { auto outputs = (*modules_[kModuleName])(input_tensors); diff --git a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc index c0369cdeb..f938c28fb 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(); +} + +ShardedStateDict PipelineParallel::BuildShardedStateDict(const std::string &prefix) const { + return modules_.at(kModuleName)->BuildShardedStateDict(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 65c55bbba..890d80d0b 100644 --- a/infini_train/src/nn/parallel/tensor_parallel.cc +++ b/infini_train/src/nn/parallel/tensor_parallel.cc @@ -284,6 +284,36 @@ bool ColumnParallelLinear::input_is_parallel() const { return input_is_parallel_ bool ColumnParallelLinear::skip_bias_add() const { return skip_bias_add_; } bool ColumnParallelLinear::sequence_parallel() const { return sequence_parallel_; } +ShardedStateDict ColumnParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {output_size_per_partition_ * tp_size, weight->Dims()[1]}; + w.local_shape = weight->Dims(); + w.global_offset = {output_size_per_partition_ * tp_rank, 0}; + w.axis_fragmentations = {tp_size, 1}; + sd.tensors[w.key] = std::move(w); + + // Bias is also split along dim=0 + if (bias_) { + auto &bias = parameter(kParamBiasName); + ShardedTensor b; + b.key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; + b.dtype = bias->Dtype(); + b.global_shape = {static_cast(output_size_per_partition_ * tp_size)}; + b.local_shape = bias->Dims(); + b.global_offset = {output_size_per_partition_ * tp_rank}; + b.axis_fragmentations = {tp_size}; + sd.tensors[b.key] = std::move(b); + } + + return sd; +} + RowParallelLinear::RowParallelLinear(int64_t in_features, int64_t out_features, bool bias, bool reduce_output, bool input_is_parallel, bool skip_bias_add, bool sequence_parallel) : CloneableModule(kType), bias_(bias), reduce_output_(reduce_output), input_is_parallel_(input_is_parallel), @@ -342,9 +372,35 @@ 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_; } +ShardedStateDict RowParallelLinear::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {weight->Dims()[0], input_size_per_partition_ * tp_size}; + w.local_shape = weight->Dims(); + w.global_offset = {0, input_size_per_partition_ * tp_rank}; + w.axis_fragmentations = {1, tp_size}; + sd.tensors[w.key] = std::move(w); + + // Bias is NOT sharded in RowParallelLinear + if (bias_) { + auto &bias = parameter(kParamBiasName); + const auto key = prefix.empty() ? kParamBiasName : prefix + "." + kParamBiasName; + auto b = MakeShardedTensor(key, bias->Dtype(), bias->Dims()); + sd.tensors[b.key] = std::move(b); + } + + return sd; +} + VocabParallelEmbedding::VocabParallelEmbedding(int64_t num_embeddings, int64_t embedding_dim, bool reduce_scatter_embeddings) - : CloneableModule(kType), embedding_dim_(embedding_dim), reduce_scatter_embeddings_(reduce_scatter_embeddings) { + : CloneableModule(kType), vocab_size_global_(num_embeddings), embedding_dim_(embedding_dim), + reduce_scatter_embeddings_(reduce_scatter_embeddings) { auto tp_size = global::GetTensorParallelSize(); CHECK_GT(tp_size, 0) << "No available devices found for VocabParallelEmbedding"; CHECK_GT(num_embeddings, 0); @@ -398,6 +454,24 @@ VocabParallelEmbedding::Forward(const std::vector> &inpu return {output}; } +ShardedStateDict VocabParallelEmbedding::BuildShardedStateDict(const std::string &prefix) const { + ShardedStateDict sd; + int tp_size = global::GetTensorParallelSize(); + + auto &weight = parameter(kParamWeightName); + ShardedTensor w; + w.key = prefix.empty() ? kParamWeightName : prefix + "." + kParamWeightName; + w.dtype = weight->Dtype(); + w.global_shape = {vocab_size_global_, embedding_dim_}; + w.local_shape = weight->Dims(); + w.global_offset = {vocab_start_index_, 0}; + w.axis_fragmentations = {tp_size, 1}; + w.allow_shape_mismatch = true; + sd.tensors[w.key] = std::move(w); + + return sd; +} + std::vector> VocabParallelCrossEntropy::Forward(const std::vector> &input_tensors) { CHECK_EQ(input_tensors.size(), 2) << kType << " expects {logits, target}"; @@ -468,7 +542,7 @@ VocabParallelCrossEntropy::Forward(const std::vector> &i auto sum_exp_local = exp_local->Sum(-1); auto sum_exp = (tp_size > 1) ? ReduceFromTPRegionFunc(sum_exp_local)[0] : sum_exp_local; - // 4. Perform Softmax(local shards but normalize globally) + // 4. Perform Softmax (local shards but normalize globally). auto softmax_local = exp_local->Div(sum_exp->Unsqueeze(-1)); // 5. Perform allreduce to get global predicted_logit @@ -482,7 +556,7 @@ VocabParallelCrossEntropy::Forward(const std::vector> &i auto log_sum_exp = sum_exp->Log(); auto loss = log_sum_exp->Sub(predicted); - // 7. Label smoothing(According to Megatron-LM) + // 7. Apply label smoothing according to Megatron-LM. // TODO(zbl): adjust smoothing coef according to vocab_size_original if (label_smoothing_ > 0.0f) { // mean_logp over *valid tokens only*: diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 39b999c77..666a0e520 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -135,29 +135,30 @@ std::unordered_map> Adam::StateDict() const std::unordered_map> state; for (size_t i = 0; i < m_.size(); ++i) { const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i]; - state.emplace("adam.m." + suffix, m_[i]); - state.emplace("adam.v." + suffix, v_[i]); + state.emplace(std::string(kAdamFirstMomentPrefix) + suffix, m_[i]); + state.emplace(std::string(kAdamSecondMomentPrefix) + suffix, v_[i]); } auto t_tensor = std::make_shared(std::vector{}, DataType::kINT64, Device()); *static_cast(t_tensor->DataPtr()) = t_; - state.emplace("adam.t", t_tensor); + state.emplace(std::string(kAdamStepKey), t_tensor); return state; } void Adam::LoadStateDict(const std::unordered_map> &state_dict) { for (size_t i = 0; i < m_.size(); ++i) { const auto suffix = parameter_names_.empty() ? std::to_string(i) : parameter_names_[i]; - const auto m_key = "adam.m." + suffix; - const auto v_key = "adam.v." + suffix; + const auto m_key = std::string(kAdamFirstMomentPrefix) + suffix; + const auto v_key = std::string(kAdamSecondMomentPrefix) + suffix; CHECK(state_dict.contains(m_key)) << "Missing optimizer state: " << m_key; CHECK(state_dict.contains(v_key)) << "Missing optimizer state: " << v_key; m_[i]->CopyFrom(state_dict.at(m_key)); v_[i]->CopyFrom(state_dict.at(v_key)); } - CHECK(state_dict.contains("adam.t")) << "Missing optimizer state: adam.t"; - const Tensor t_cpu = state_dict.at("adam.t")->To(Device()); + const std::string step_key(kAdamStepKey); + CHECK(state_dict.contains(step_key)) << "Missing optimizer state: " << step_key; + const Tensor t_cpu = state_dict.at(step_key)->To(Device()); t_ = *static_cast(t_cpu.DataPtr()); } } // namespace optimizers diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index 34edd26ff..84c6351f6 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -374,9 +374,10 @@ args_string_for_test() { jq -r --argjson g "$group_idx" --argjson t "$test_idx" --arg model "$model_name" --arg test_id "$test_id" ' def namespaced_path($p; $model; $mode): - if ($p | test("/checkpoint_step_[0-9]+($|/)")) then - ($p | capture("^(?.*)/(?checkpoint_step_[0-9]+(?:/.*)?)$")) as $m - | ($m.prefix + "/" + $model + "/" + $mode + "/" + $m.step) + if ($p | test("/(?:checkpoint_step_|iter_)[0-9]+($|/)")) then + ($p | capture("^(?.*)/(?:checkpoint_step_|iter_)(?[0-9]+)(?/.*)?$")) as $m + | ("0000000" + $m.iteration)[-7:] as $iteration + | ($m.prefix + "/" + $model + "/" + $mode + "/iter_" + $iteration + ($m.suffix // "")) else ($p + "/" + $model + "/" + $mode) end; diff --git a/tests/checkpoint/test_checkpoint_serialization.cc b/tests/checkpoint/test_checkpoint_serialization.cc index 8a9b11cf4..59d6ab3bd 100644 --- a/tests/checkpoint/test_checkpoint_serialization.cc +++ b/tests/checkpoint/test_checkpoint_serialization.cc @@ -1,12 +1,22 @@ +#include #include +#include +#include +#include #include #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/constants.h" +#include "infini_train/include/checkpoint/load_planner.h" +#include "infini_train/include/checkpoint/load_strategy.h" +#include "infini_train/include/checkpoint/save_planner.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/modules/transformer/mlp.h" #include "infini_train/include/optimizer.h" +#include "infini_train/include/shard_spec.h" #include "infini_train/include/tensor.h" #include "tests/common/test_utils.h" @@ -14,8 +24,76 @@ using namespace infini_train; namespace nn = infini_train::nn; +namespace { + +class NamedParameterModule final : public nn::Module { +public: + void AddParameter(const std::string &name, const std::shared_ptr ¶meter) { + parameters_[name] = parameter; + } + + void AddModule(const std::string &name, const std::shared_ptr &module) { modules_[name] = module; } +}; + +} // namespace + +TEST(ModuleNamedParametersTest, SupportsTorchStyleArgumentsAndSharedParameterDeduplication) { + auto root = std::make_shared(); + auto child = std::make_shared(); + auto shared = std::make_shared(std::vector{1}, DataType::kFLOAT32, Device()); + auto child_weight = std::make_shared(std::vector{1}, DataType::kFLOAT32, Device()); + root->AddParameter("root_weight", shared); + child->AddParameter("alias", shared); + child->AddParameter("weight", child_weight); + root->AddModule("child", child); + + const auto local = root->NamedParameters("model", false); + ASSERT_EQ(local.size(), 1); + EXPECT_EQ(local[0].first, "model.root_weight"); + EXPECT_EQ(local[0].second, shared); + + const auto deduplicated = root->NamedParameters("model"); + ASSERT_EQ(deduplicated.size(), 2); + EXPECT_EQ(deduplicated[0].first, "model.root_weight"); + EXPECT_EQ(deduplicated[1].first, "model.child.weight"); + + const auto aliases = root->NamedParameters("model", true, false); + ASSERT_EQ(aliases.size(), 3); + EXPECT_EQ(aliases[0].first, "model.root_weight"); + EXPECT_EQ(aliases[1].first, "model.child.alias"); + EXPECT_EQ(aliases[1].second, shared); + EXPECT_EQ(aliases[2].first, "model.child.weight"); +} + class CheckpointSerializationTest : public test::InfiniTrainTest {}; +TEST(ShardedStateDictTest, RejectsDuplicateKeysWhenMerging) { + ShardedStateDict destination; + destination.tensors["weight"] = {.key = "weight"}; + ShardedStateDict source; + source.tensors["weight"] = {.key = "weight"}; + + EXPECT_DEATH(destination.Merge(std::move(source)), "Duplicate sharded state-dict key: weight"); +} + +TEST(ShardedStateDictTest, PackedSwiGLUUsesGateAndUpSegments) { + nn::TransformerConfig config; + config.n_embd = 4; + config.activation_type = nn::MLPType::kSwiGLU; + config.ffn_expansion_ratio = 3.0f; + config.ffn_dim_multiplier = std::nullopt; + config.multiple_of = 1; + + const auto state = nn::MLP(config).BuildShardedStateDict("mlp"); + const auto &weight = state.tensors.at("mlp.c_fc.weight"); + ASSERT_EQ(weight.segments.size(), 2); + EXPECT_EQ(weight.segments[0], (ShardSegment{.global_offset = 0, .local_offset = 0, .length = 8})); + EXPECT_EQ(weight.segments[1], (ShardSegment{.global_offset = 8, .local_offset = 8, .length = 8})); + + const auto &bias = state.tensors.at("mlp.c_fc.bias"); + EXPECT_EQ(bias.segments, weight.segments); +} + TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { auto dir = std::filesystem::temp_directory_path() / "test_ckpt_fp32"; std::filesystem::remove_all(dir); @@ -28,7 +106,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { p2->Fill(-1.5f); *model1->mutable_parameter("bias") = p2; - auto opt1 = std::make_shared(model1->Parameters(), 0.01); + auto opt1 = optimizers::Adam::CreateNamed(0.01)(model1->NamedParameters()); TrainerState saved{.global_step = 42, .consumed_train_samples = 100}; Checkpoint::Save(dir, *model1, opt1.get(), saved, nullptr); @@ -39,7 +117,7 @@ TEST_P(CheckpointSerializationTest, SaveAndLoadModelFP32) { auto q2 = std::make_shared(std::vector{4}, DataType::kFLOAT32, GetDevice()); q2->Fill(0.0f); *model2->mutable_parameter("bias") = q2; - auto opt2 = std::make_shared(model2->Parameters(), 0.01); + auto opt2 = optimizers::Adam::CreateNamed(0.01)(model2->NamedParameters()); TrainerState loaded; Checkpoint::Load(dir, *model2, opt2.get(), loaded, nullptr); @@ -52,4 +130,559 @@ 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 / checkpoint::kModelCheckpointFilename; + Checkpoint::SaveStateDictFile(path, {{"matrix", matrix}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("matrix") - 1 + sizeof(int8_t) + + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + 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 = checkpoint::kModelCheckpointFilename, + .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_P(CheckpointSerializationTest, ConvertsSavedBF16TensorToFP32Target) { + auto dir = std::filesystem::temp_directory_path() / "test_ckpt_bf16_to_fp32"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + + auto source_fp32 = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, Device()); + auto *source_data = static_cast(source_fp32->DataPtr()); + source_data[0] = 1.0f; + source_data[1] = 2.0f; + source_data[2] = 3.0f; + source_data[3] = 4.0f; + auto source_bf16 = std::make_shared(source_fp32->To(DataType::kBFLOAT16)); + Checkpoint::SaveStateDictFile(dir / checkpoint::kModelCheckpointFilename, {{"weight", source_bf16}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("weight") - 1 + sizeof(int8_t) + + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + + checkpoint::LoadPlan plan; + plan.tensors["weight"] = {.key = "weight", + .dtype = DataType::kFLOAT32, + .global_shape = {2, 2}, + .target_shape = {2, 2}, + .reads = {{.key = "weight", + .filename = checkpoint::kModelCheckpointFilename, + .dtype = DataType::kBFLOAT16, + .global_shape = {2, 2}, + .byte_size = source_bf16->SizeInBytes(), + .data_offset = data_offset, + .shard_dim = -1, + .source_shape = {2, 2}}}}; + + checkpoint::IndexedRegionLoadStrategy strategy; + const auto loaded = strategy.Execute(dir, plan).at("weight"); + ASSERT_EQ(loaded->Dtype(), DataType::kFLOAT32); + const auto *loaded_data = static_cast(loaded->DataPtr()); + for (int i = 0; i < 4; ++i) { EXPECT_FLOAT_EQ(loaded_data[i], source_data[i]); } + + std::filesystem::remove_all(dir); +} + +TEST(CheckpointLoadPlannerTest, PadsVocabularyTailWhenTargetTpUsesPaddedVocab) { + const auto dir = std::filesystem::temp_directory_path() / "test_vocab_padding_reshard"; + std::filesystem::remove_all(dir); + 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 / checkpoint::kModelCheckpointFilename, {{"lm_head.weight", source}}); + constexpr uint64_t data_offset = sizeof(uint32_t) * 3 + sizeof(uint32_t) + sizeof("lm_head.weight") - 1 + + sizeof(int8_t) + sizeof(uint32_t) + sizeof(int64_t) * 2 + sizeof(uint64_t); + + 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 = checkpoint::kModelCheckpointFilename, + .offset = data_offset, + .byte_size = sizeof(float) * 10}); + + 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}, + .allow_shape_mismatch = true}; + + 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(CheckpointMetadataTest, RejectsMalformedIntegerArrays) { + const auto dir = std::filesystem::temp_directory_path() / "test_malformed_checkpoint_metadata"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + { + std::ofstream metadata_file(dir / checkpoint::kMetadataFilename); + metadata_file << R"({ + "version": 3, + "format": "infinitrain_sharded", + "tensors": [ + { + "key": "weight", + "dtype": "float32", + "global_shape": [8, invalid] + } + ] +})"; + } + + EXPECT_DEATH(Checkpoint::LoadMetadata(dir), "Invalid integer in checkpoint metadata field global_shape"); + std::filesystem::remove_all(dir); +} + +TEST_P(CheckpointSerializationTest, GlobalMetadataRoundTrip) { + auto dir = std::filesystem::temp_directory_path() / "test_global_metadata"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(dir); + Checkpoint::CheckpointMetadata metadata; + metadata.version = 3; + metadata.has_metadata = true; + 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 / checkpoint::kMetadataFilename, metadata); + std::ifstream metadata_file(dir / checkpoint::kMetadataFilename); + const std::string metadata_json((std::istreambuf_iterator(metadata_file)), std::istreambuf_iterator()); + EXPECT_EQ(metadata_json.find("\"iteration\""), std::string::npos); + EXPECT_EQ(metadata_json.find("\"parallel_config\""), std::string::npos); + EXPECT_EQ(metadata_json.find("\"model_config\""), std::string::npos); + + auto loaded = Checkpoint::LoadMetadata(dir); + ASSERT_TRUE(loaded.has_metadata); + 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], (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}; +} + +ShardedStateDict MakeTarget(const std::string &key, int count, int index, int64_t global_size) { + ShardedStateDict target; + target.tensors[key] = {.key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {global_size, 4}, + .local_shape = {global_size / count, 4}, + .global_offset = {global_size / count * index, 0}, + .axis_fragmentations = {count, 1}}; + return target; +} + +std::shared_ptr MakeColumnTensor(const std::vector &values) { + auto tensor = std::make_shared(std::vector{static_cast(values.size()), 1}, + DataType::kFLOAT32, Device()); + std::copy(values.begin(), values.end(), static_cast(tensor->DataPtr())); + return tensor; +} + +uint64_t SingleTensorDataOffset(const std::string &key, size_t dimensions) { + return sizeof(uint32_t) * 3 + sizeof(uint32_t) + key.size() + sizeof(int8_t) + sizeof(uint32_t) + + sizeof(int64_t) * dimensions + sizeof(uint64_t); +} +} // namespace + +TEST(CheckpointOptimizerShardingTest, AdamMomentsReuseModelShardMetadata) { + 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"); + EXPECT_EQ(m.dtype, moment->Dtype()); + EXPECT_EQ(m.allow_shape_mismatch, model.tensors.at("c_attn.weight").allow_shape_mismatch); + const auto &t = optimizer.tensors.at("adam.t"); + EXPECT_TRUE(t.global_shape.empty()); + EXPECT_TRUE(t.local_shape.empty()); +} + +TEST(CheckpointOptimizerShardingTest, AdamMomentUsesOptimizerStateDtype) { + ShardedStateDict model; + model.tensors["weight"] = {.key = "weight", + .dtype = DataType::kBFLOAT16, + .global_shape = {4, 4}, + .local_shape = {4, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}}; + auto moment = std::make_shared(std::vector{4, 4}, DataType::kFLOAT32, Device()); + + const auto optimizer + = checkpoint::BuildOptimizerShardedStateDict(model, {{"adam.m.weight", moment}, {"adam.v.weight", moment}}); + EXPECT_EQ(optimizer.tensors.at("adam.m.weight").dtype, DataType::kFLOAT32); + EXPECT_EQ(optimizer.tensors.at("adam.v.weight").dtype, DataType::kFLOAT32); +} + +TEST(CheckpointLoadPlannerTest, RejectsUnknownCheckpointDtype) { + auto metadata = Checkpoint::CheckpointMetadata{}; + metadata.tensors = {MakeSavedShard("weight", 1, 0, 16, "rank_0/model.ckpt")}; + metadata.tensors.front().dtype_str = "unknown_dtype"; + + EXPECT_DEATH(checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 1, 0, 16)), + "Unsupported checkpoint tensor dtype: unknown_dtype"); +} + +TEST(CheckpointLoadPlannerTest, TensorParallelTwoToFourReadsOnlyOverlap) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {MakeSavedShard("weight", 2, 0, 16, "rank_0/model.ckpt"), + MakeSavedShard("weight", 2, 1, 16, "rank_1/model.ckpt")}; + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 4, 1, 16)); + const auto &reads = plan.tensors.at("weight").reads; + ASSERT_EQ(reads.size(), 1); + EXPECT_EQ(reads[0].filename, "rank_0/model.ckpt"); + EXPECT_EQ(reads[0].source_offset, 4); + EXPECT_EQ(reads[0].target_offset, 0); + EXPECT_EQ(reads[0].length, 4); +} + +TEST(CheckpointLoadPlannerTest, TensorParallelFourToTwoReadsTwoOverlaps) { + Checkpoint::CheckpointMetadata metadata; + for (int index = 0; index < 4; ++index) { + metadata.tensors.push_back( + MakeSavedShard("weight", 4, index, 16, "rank_" + std::to_string(index) + "/model.ckpt")); + } + + auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, MakeTarget("weight", 2, 1, 16)); + const auto &reads = plan.tensors.at("weight").reads; + ASSERT_EQ(reads.size(), 2); + EXPECT_EQ(reads[0].filename, "rank_2/model.ckpt"); + EXPECT_EQ(reads[0].target_offset, 0); + EXPECT_EQ(reads[1].filename, "rank_3/model.ckpt"); + EXPECT_EQ(reads[1].target_offset, 4); +} + +TEST(CheckpointLoadPlannerTest, UsesExplicitGlobalOffsetsForUnevenShards) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {{.key = "weight", + .dtype_str = "float32", + .global_shape = {8, 4}, + .local_shape = {3, 4}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .file = "rank_0/model.ckpt"}, + {.key = "weight", + .dtype_str = "float32", + .global_shape = {8, 4}, + .local_shape = {5, 4}, + .global_offset = {3, 0}, + .axis_fragmentations = {2, 1}, + .file = "rank_1/model.ckpt"}}; + 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)); + + 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)); + } + + 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, PackedSwiGLUExecutesTensorParallelTwoToOne) { + const auto dir = std::filesystem::temp_directory_path() / "test_swiglu_tp2_to_tp1"; + std::filesystem::remove_all(dir); + const std::string key = "c_fc.weight"; + const auto data_offset = SingleTensorDataOffset(key, 2); + + Checkpoint::CheckpointMetadata metadata; + for (int rank = 0; rank < 2; ++rank) { + const auto rank_dir = dir / ("rank_" + std::to_string(rank)); + std::filesystem::create_directories(rank_dir); + const std::vector values = rank == 0 ? std::vector{0, 1, 4, 5} : std::vector{2, 3, 6, 7}; + const auto tensor = MakeColumnTensor(values); + Checkpoint::SaveStateDictFile(rank_dir / checkpoint::kModelCheckpointFilename, {{key, tensor}}); + + metadata.tensors.push_back({ + .key = key, + .dtype_str = "fp32", + .global_shape = {8, 1}, + .local_shape = {4, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = { + {.global_offset = rank * 2, .local_offset = 0, .length = 2}, + {.global_offset = 4 + rank * 2, .local_offset = 2, .length = 2}, + }, + .file = "rank_" + std::to_string(rank) + "/" + checkpoint::kModelCheckpointFilename, + .offset = data_offset, + .byte_size = tensor->SizeInBytes(), + }); + } + + ShardedStateDict target; + target.tensors[key] = { + .key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {8, 1}, + .local_shape = {8, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 4}, + {.global_offset = 4, .local_offset = 4, .length = 4}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto loaded = checkpoint::IndexedRegionLoadStrategy().Execute(dir, plan).at(key); + const auto *data = static_cast(loaded->DataPtr()); + for (int row = 0; row < 8; ++row) { EXPECT_FLOAT_EQ(data[row], static_cast(row)); } + std::filesystem::remove_all(dir); +} + +TEST(CheckpointLoadPlannerTest, PackedSwiGLUExecutesTensorParallelOneToTwo) { + const auto dir = std::filesystem::temp_directory_path() / "test_swiglu_tp1_to_tp2"; + const auto rank_dir = dir / "rank_0"; + std::filesystem::remove_all(dir); + std::filesystem::create_directories(rank_dir); + const std::string key = "c_fc.weight"; + const auto tensor = MakeColumnTensor({0, 1, 2, 3, 4, 5, 6, 7}); + Checkpoint::SaveStateDictFile(rank_dir / checkpoint::kModelCheckpointFilename, {{key, tensor}}); + + Checkpoint::CheckpointMetadata metadata; + metadata.tensors.push_back({ + .key = key, + .dtype_str = "fp32", + .global_shape = {8, 1}, + .local_shape = {8, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {1, 1}, + .segments = { + {.global_offset = 0, .local_offset = 0, .length = 4}, + {.global_offset = 4, .local_offset = 4, .length = 4}, + }, + .file = "rank_0/" + std::string(checkpoint::kModelCheckpointFilename), + .offset = SingleTensorDataOffset(key, 2), + .byte_size = tensor->SizeInBytes(), + }); + + ShardedStateDict target; + target.tensors[key] = { + .key = key, + .dtype = DataType::kFLOAT32, + .global_shape = {8, 1}, + .local_shape = {4, 1}, + .global_offset = {0, 0}, + .axis_fragmentations = {2, 1}, + .segments = { + {.global_offset = 2, .local_offset = 0, .length = 2}, + {.global_offset = 6, .local_offset = 2, .length = 2}, + }, + }; + + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + const auto loaded = checkpoint::IndexedRegionLoadStrategy().Execute(dir, plan).at(key); + const auto *data = static_cast(loaded->DataPtr()); + const std::vector expected = {2, 3, 6, 7}; + for (size_t row = 0; row < expected.size(); ++row) { EXPECT_FLOAT_EQ(data[row], expected[row]); } + std::filesystem::remove_all(dir); +} + +TEST(CheckpointLoadPlannerTest, PipelineReshardPlansOnlyTargetStageKeys) { + Checkpoint::CheckpointMetadata metadata; + metadata.tensors = {MakeSavedShard("layer.0.weight", 1, 0, 8, "old_pp0/model.ckpt"), + 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"); +} + +TEST(CheckpointLoadPlannerTest, VirtualPipelineReshardUsesGlobalLayerKeys) { + Checkpoint::CheckpointMetadata metadata; + for (int layer = 0; layer < 4; ++layer) { + const auto key = "transformer.h." + std::to_string(layer) + ".weight"; + metadata.tensors.push_back( + MakeSavedShard(key, 1, 0, 8, "saved_layer_" + std::to_string(layer) + "/model.ckpt")); + } + + auto target = MakeTarget("transformer.h.0.weight", 1, 0, 8); + target.Merge(MakeTarget("transformer.h.2.weight", 1, 0, 8)); + const auto plan = checkpoint::LoadPlanner::PlanReshard(metadata, target); + + ASSERT_EQ(plan.tensors.size(), 2); + EXPECT_EQ(plan.tensors.at("transformer.h.0.weight").reads[0].filename, "saved_layer_0/model.ckpt"); + EXPECT_EQ(plan.tensors.at("transformer.h.2.weight").reads[0].filename, "saved_layer_2/model.ckpt"); +} diff --git a/tests/checkpoint/test_lr_scheduler_state.cc b/tests/checkpoint/test_lr_scheduler_state.cc index a0e07dd93..40817b2d9 100644 --- a/tests/checkpoint/test_lr_scheduler_state.cc +++ b/tests/checkpoint/test_lr_scheduler_state.cc @@ -5,6 +5,7 @@ #include "gtest/gtest.h" #include "infini_train/include/checkpoint/checkpoint.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/lr_scheduler.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/optimizer.h" @@ -61,7 +62,7 @@ TEST_P(LRSchedulerCheckpointTest, SaveAndLoadLRSchedulerState) { TrainerState saved{.global_step = 3, .consumed_train_samples = 12}; Checkpoint::Save(dir, *model1, nullptr, saved, sched1.get()); - EXPECT_TRUE(std::filesystem::exists(dir / "lr_scheduler.ckpt")); + EXPECT_TRUE(std::filesystem::exists(dir / checkpoint::kLRSchedulerFilename)); auto model2 = MakeModel(GetDevice()); auto opt2 = std::make_shared(model2->Parameters(), kBaseLR); @@ -92,7 +93,7 @@ TEST_P(LRSchedulerCheckpointTest, SkipsLRSchedulerStateWhenSchedulerIsNull) { TrainerState saved{.global_step = 3}; Checkpoint::Save(dir, *model1, nullptr, saved, nullptr); - EXPECT_FALSE(std::filesystem::exists(dir / "lr_scheduler.ckpt")); + EXPECT_FALSE(std::filesystem::exists(dir / checkpoint::kLRSchedulerFilename)); std::filesystem::remove_all(dir); } diff --git a/tests/checkpoint/test_optimizer_state.cc b/tests/checkpoint/test_optimizer_state.cc index 1cbb8b9f3..aa61b3354 100644 --- a/tests/checkpoint/test_optimizer_state.cc +++ b/tests/checkpoint/test_optimizer_state.cc @@ -78,6 +78,24 @@ 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()); + const NamedParameterList named_parameters{{"transformer.h.0.weight", first}, {"transformer.h.0.bias", second}}; + auto adam = std::make_shared(named_parameters, 0.001); + + const auto state = adam->StateDict(); + EXPECT_TRUE(state.contains("adam.m.transformer.h.0.weight")); + 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(named_parameters, 0.001); + 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/checkpoint/test_trainer_state.cc b/tests/checkpoint/test_trainer_state.cc index ec4d61e8b..6df573892 100644 --- a/tests/checkpoint/test_trainer_state.cc +++ b/tests/checkpoint/test_trainer_state.cc @@ -6,8 +6,10 @@ #include "infini_train/include/checkpoint/checkpoint.h" #include "infini_train/include/checkpoint/checkpoint_manager.h" +#include "infini_train/include/checkpoint/constants.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -26,11 +28,13 @@ TEST_P(TrainerStateTest, DefaultValues) { EXPECT_EQ(state.n_head, 0); EXPECT_EQ(state.n_kv_head, 0); EXPECT_EQ(state.n_embd, 0); - EXPECT_EQ(state.vocab_size, 0); + EXPECT_EQ(state.original_vocab_size, 0); + EXPECT_EQ(state.padded_vocab_size, 0); EXPECT_EQ(state.ddp_size, 1); EXPECT_EQ(state.tp_size, 1); EXPECT_EQ(state.sp_size, 1); EXPECT_EQ(state.pp_size, 1); + EXPECT_EQ(state.vpp_size, 1); } TEST_P(TrainerStateTest, TrainerStateFileCreated) { @@ -47,9 +51,9 @@ TEST_P(TrainerStateTest, TrainerStateFileCreated) { Checkpoint::Save(dir, *model, opt.get(), saved, nullptr); - EXPECT_TRUE(std::filesystem::exists(dir / "trainer_state.json")); + EXPECT_TRUE(std::filesystem::exists(dir / checkpoint::kTrainerStateFilename)); - std::ifstream ifs(dir / "trainer_state.json"); + std::ifstream ifs(dir / checkpoint::kTrainerStateFilename); std::string content((std::istreambuf_iterator(ifs)), std::istreambuf_iterator()); EXPECT_NE(content.find("\"global_step\""), std::string::npos); EXPECT_NE(content.find("\"consumed_train_samples\""), std::string::npos); @@ -68,11 +72,13 @@ TEST_P(TrainerStateTest, RoundTrip) { .n_head = 16, .n_kv_head = 8, .n_embd = 1024, - .vocab_size = 128256, + .original_vocab_size = 128000, + .padded_vocab_size = 128256, .ddp_size = 2, .tp_size = 1, .sp_size = 1, .pp_size = 2, + .vpp_size = 2, }; auto model1 = std::make_shared(1, 3, true, GetDevice()); @@ -98,9 +104,14 @@ TEST_P(TrainerStateTest, RoundTrip) { EXPECT_EQ(loaded.n_head, 16); EXPECT_EQ(loaded.n_kv_head, 8); EXPECT_EQ(loaded.n_embd, 1024); - EXPECT_EQ(loaded.vocab_size, 128256); - EXPECT_EQ(loaded.ddp_size, 2); - EXPECT_EQ(loaded.pp_size, 2); + EXPECT_EQ(loaded.original_vocab_size, 128000); + EXPECT_EQ(loaded.padded_vocab_size, 128256); + EXPECT_EQ(loaded.ddp_size, nn::parallel::global::GetDataParallelSize()); + EXPECT_EQ(loaded.tp_size, nn::parallel::global::GetTensorParallelSize()); + EXPECT_EQ(loaded.pp_size, nn::parallel::global::GetPipelineParallelSize()); + EXPECT_EQ(loaded.sp_size, + nn::parallel::global::GetSequenceParallelEnabled() ? nn::parallel::global::GetTensorParallelSize() : 1); + EXPECT_EQ(loaded.vpp_size, nn::parallel::global::GetVirtualPipelineParallelSize()); std::filesystem::remove_all(dir); } diff --git a/tests/lora/test_lora.cc b/tests/lora/test_lora.cc index 26cffdcaa..01e2725a2 100644 --- a/tests/lora/test_lora.cc +++ b/tests/lora/test_lora.cc @@ -7,11 +7,13 @@ #include "infini_train/include/nn/lora/lora_config.h" #include "infini_train/include/nn/lora/lora_linear.h" +#include "infini_train/include/nn/lora/lora_parallel_linear.h" #include "infini_train/include/nn/lora/lora_utils.h" #include "infini_train/include/nn/modules/container.h" #include "infini_train/include/nn/modules/linear.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/tensor_parallel.h" #include "infini_train/include/tensor.h" #include "tests/common/test_utils.h" @@ -78,6 +80,33 @@ TEST_P(LoRATest, LoRAConfigScaling) { EXPECT_EQ(config.Scaling(), expected_scaling); } +TEST_P(LoRATest, ParallelLoRAShardedStateDictIncludesAdapterParameters) { + LoRAConfig config; + config.rank = 2; + + auto column_base = std::make_shared( + 4, 6, /*bias=*/false, /*gather_output=*/false, /*input_is_parallel=*/false, /*skip_bias_add=*/false, + /*sequence_parallel=*/false); + auto column = std::make_shared(column_base, config, 4, 6); + const auto column_state = column->BuildShardedStateDict("column"); + ASSERT_TRUE(column_state.tensors.contains("column.lora_A")); + ASSERT_TRUE(column_state.tensors.contains("column.lora_B")); + EXPECT_EQ(column_state.tensors.at("column.lora_A").axis_fragmentations, (std::vector{1, 1})); + EXPECT_EQ(column_state.tensors.at("column.lora_B").axis_fragmentations, + (std::vector{nn::parallel::global::GetTensorParallelSize(), 1})); + + auto row_base = std::make_shared( + 4, 6, /*bias=*/false, /*reduce_output=*/true, /*input_is_parallel=*/true, /*skip_bias_add=*/false, + /*sequence_parallel=*/false); + auto row = std::make_shared(row_base, config, 4, 6); + const auto row_state = row->BuildShardedStateDict("row"); + ASSERT_TRUE(row_state.tensors.contains("row.lora_A")); + ASSERT_TRUE(row_state.tensors.contains("row.lora_B")); + EXPECT_EQ(row_state.tensors.at("row.lora_A").axis_fragmentations, + (std::vector{1, nn::parallel::global::GetTensorParallelSize()})); + EXPECT_EQ(row_state.tensors.at("row.lora_B").axis_fragmentations, (std::vector{1, 1})); +} + TEST_P(LoRATest, PackedQKVShardGPTStyle) { auto full_qkv = MakeRowLabeledTensor(/*rows=*/12, /*cols=*/3, GetDevice()); auto shard = infini_train::nn::lora::detail::SlicePackedQKVRowsForTensorParallel(full_qkv, /*q_rows=*/4, diff --git a/tests/transformer/test_transformer_architecture.cc b/tests/transformer/test_transformer_architecture.cc index 4cec471de..f6ed318f6 100644 --- a/tests/transformer/test_transformer_architecture.cc +++ b/tests/transformer/test_transformer_architecture.cc @@ -163,6 +163,10 @@ TEST_P(TransformerModuleTest, LLaMA3Model) { auto model = std::make_shared(config); model->To(GetDevice()); EXPECT_FALSE(model->Parameters().empty()); + const auto sharded_state = model->BuildShardedStateDict(); + 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) {