From ffb5daba2b70f40630d17d7c7d33d2aa5d3c3ef4 Mon Sep 17 00:00:00 2001 From: why1te Date: Sat, 19 Sep 2026 04:35:56 +0000 Subject: [PATCH] feat: support custom pipeline layouts - add a shared PipelineLayout model for uniform partitions, uneven layer partitions, and explicit chunk-to-stage ownership - construct GPT-2 and LLaMA models, checkpoint loaders, and pipeline schedules from the resolved layout while preserving default behavior and legacy API compatibility - handle logical input/output stages, local chunk chains, and cyclic owner mappings without changing legacy schedule ordering - add parser, layout, loader, wrapper, and scheduler regression tests --- CMakeLists.txt | 2 + example/common/parser.cc | 212 ++++++++++ example/common/parser.h | 34 ++ example/gpt2/checkpoint_loader.cc | 157 ++++--- example/gpt2/checkpoint_loader.h | 5 + example/gpt2/main.cc | 100 +++-- example/llama3/checkpoint_loader.cc | 170 +++++--- example/llama3/checkpoint_loader.h | 5 + example/llama3/main.cc | 85 +++- .../nn/modules/transformer/transformer.h | 11 + .../include/nn/parallel/pp/pipeline_layout.h | 62 +++ .../nn/parallel/pp/pipeline_parallel.h | 8 + .../nn/parallel/pp/pipeline_schedule.h | 17 +- .../include/utils/precision_checker.h | 6 +- .../src/nn/modules/transformer/transformer.cc | 80 +++- .../src/nn/parallel/pp/pipeline_layout.cc | 260 ++++++++++++ .../src/nn/parallel/pp/pipeline_parallel.cc | 112 +++-- .../src/nn/parallel/pp/pipeline_schedule.cc | 323 ++++++++++----- infini_train/src/utils/precision_checker.cc | 53 ++- tests/distributed/CMakeLists.txt | 41 ++ tests/distributed/test_pipeline_layout.cc | 315 ++++++++++++++ tests/distributed/test_pipeline_loader.cc | 241 +++++++++++ tests/distributed/test_pipeline_parallel.cc | 221 ++++++++++ tests/distributed/test_pipeline_schedule.cc | 389 ++++++++++++++++++ tests/distributed/test_pp_layout_parser.cc | 115 ++++++ 25 files changed, 2714 insertions(+), 310 deletions(-) create mode 100644 example/common/parser.cc create mode 100644 example/common/parser.h create mode 100644 infini_train/include/nn/parallel/pp/pipeline_layout.h create mode 100644 infini_train/src/nn/parallel/pp/pipeline_layout.cc create mode 100644 tests/distributed/test_pipeline_layout.cc create mode 100644 tests/distributed/test_pipeline_loader.cc create mode 100644 tests/distributed/test_pipeline_parallel.cc create mode 100644 tests/distributed/test_pipeline_schedule.cc create mode 100644 tests/distributed/test_pp_layout_parser.cc diff --git a/CMakeLists.txt b/CMakeLists.txt index 4ffbc25eb..397466dfd 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -199,6 +199,7 @@ add_executable(gpt2 example/gpt2/main.cc example/common/tiny_shakespeare_dataset.cc example/common/utils.cc + example/common/parser.cc example/common/tokenizer.cc example/gpt2/checkpoint_loader.cc ) @@ -215,6 +216,7 @@ link_infini_train_exe(mixtral) add_executable(llama3 example/llama3/main.cc example/common/tiny_shakespeare_dataset.cc + example/common/parser.cc example/common/utils.cc example/common/tokenizer.cc example/llama3/checkpoint_loader.cc diff --git a/example/common/parser.cc b/example/common/parser.cc new file mode 100644 index 000000000..434c122ef --- /dev/null +++ b/example/common/parser.cc @@ -0,0 +1,212 @@ +#include "example/common/parser.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace infini_train::examples { + +using LayerIndex = nn::parallel::LayerIndex; +using PipelineLayout = nn::parallel::PipelineLayout; +namespace { +std::invalid_argument InvalidPartition(std::string_view value, std::string_view reason) { + return std::invalid_argument("Invalid --pipeline_layer_partition=\"" + std::string(value) + + "\": " + std::string(reason)); +} + +} // namespace +std::vector ParsePipelineLayerPartition(std::string_view value) { + if (value.empty()) { + return {}; + } + + std::vector result; + std::size_t begin = 0; + std::size_t item_index = 0; + + while (begin <= value.size()) { + const std::size_t comma = value.find(',', begin); + const std::size_t end = comma == std::string_view::npos ? value.size() : comma; + std::string_view token = value.substr(begin, end - begin); + + if (token.empty()) { + throw InvalidPartition(value, "item " + std::to_string(item_index) + " is empty"); + } + + for (char c : token) { + if (c < '0' || c > '9') { + throw InvalidPartition(value, + "item " + std::to_string(item_index) + " must contain decimal digits only"); + } + } + LayerIndex count = 0; + const auto [ptr, error] = std::from_chars(token.data(), token.data() + token.size(), count); + + if (error == std::errc::result_out_of_range) { + throw InvalidPartition(value, "item " + std::to_string(item_index) + " is outside the LayerIndex range"); + } + if (error != std::errc{} || ptr != token.data() + token.size()) { + throw InvalidPartition(value, "item " + std::to_string(item_index) + + " must be a base-10 integer without trailing characters"); + } + if (count <= 0) { + throw InvalidPartition(value, "item " + std::to_string(item_index) + " must be greater than zero"); + } + + result.push_back(count); + if (comma == std::string_view::npos) { + break; + } + begin = comma + 1; + ++item_index; + } + + return result; +} + +std::shared_ptr BuildPipelineLayoutFromCLI(LayerIndex num_layers, int num_stages, + std::span counts, + int chunks_per_stage) { + if (num_stages <= 0 || chunks_per_stage <= 0 + || static_cast(num_stages) * chunks_per_stage > std::numeric_limits::max()) { + throw std::invalid_argument("PP/vPP must be positive and their product must fit int"); + } + if (counts.empty() && num_layers >= 0 + && num_layers < static_cast(num_stages) * chunks_per_stage) { + return nullptr; + } + PipelineLayout layout = counts.empty() + ? PipelineLayout::BuildUniformLayout(num_layers, num_stages, chunks_per_stage) + : PipelineLayout::BuildCustomLayout(num_layers, num_stages, counts, chunks_per_stage); + return std::make_shared(std::move(layout)); +} + +bool IsExplicitChunkRequest(const PipelineLayoutRequest &request) { + return std::holds_alternative>(request); +} + +PipelineLayoutRequest ParsePipelineLayoutRequest(std::string_view partition, std::string_view chunks, int pp_size, + int vpp_size) { + if (pp_size <= 0 || vpp_size <= 0 + || static_cast(pp_size) * vpp_size > std::numeric_limits::max()) { + throw std::invalid_argument("PP/vPP must be positive and their product must fit int"); + } + if (!partition.empty() && !chunks.empty()) { + throw std::invalid_argument("pipeline_layer_partition and pipeline_chunk_layout are mutually exclusive"); + } + if (!chunks.empty()) { + std::vector specs; + std::size_t begin = 0; + while (begin <= chunks.size()) { + const auto comma = chunks.find(',', begin); + const auto token + = chunks.substr(begin, comma == std::string_view::npos ? chunks.size() - begin : comma - begin); + const auto colon = token.find(':'); + if (colon == std::string_view::npos || token.find(':', colon + 1) != std::string_view::npos) { + throw std::invalid_argument("pipeline_chunk_layout entry " + std::to_string(specs.size()) + + " must be stage:count: " + std::string(token)); + } + const auto owner_text = token.substr(0, colon); + const auto count_text = token.substr(colon + 1); + const auto digits = [](std::string_view value) { + return !value.empty() && value.find_first_not_of("0123456789") == std::string_view::npos; + }; + int owner = -1; + LayerIndex count = -1; + if (!digits(owner_text) || !digits(count_text)) { + throw std::invalid_argument("invalid pipeline_chunk_layout entry " + std::to_string(specs.size()) + ": " + + std::string(token)); + } + const auto a = std::from_chars(owner_text.data(), owner_text.data() + owner_text.size(), owner); + const auto b = std::from_chars(count_text.data(), count_text.data() + count_text.size(), count); + if (a.ec != std::errc{} || b.ec != std::errc{} || owner < 0 || owner >= pp_size || count <= 0) { + throw std::invalid_argument("pipeline_chunk_layout entry " + std::to_string(specs.size()) + + " has invalid owner/count: " + std::string(token)); + } + specs.push_back({owner, count}); + if (specs.size() > static_cast(std::numeric_limits::max())) { + throw std::invalid_argument("too many explicit chunks"); + } + if (comma == std::string_view::npos) { + break; + } + begin = comma + 1; + } + std::vector local_counts(static_cast(pp_size)); + for (const auto &spec : specs) { ++local_counts[spec.stage_id]; } + int max_count = 0; + for (int count : local_counts) { + if (count == 0) { + throw std::invalid_argument("each stage must own a chunk"); + } + if (count > max_count) { + max_count = count; + } + } + if (max_count != vpp_size) { + throw std::invalid_argument("explicit layout max local chunk count=" + std::to_string(max_count) + + ", virtual_pipeline_parallel=" + std::to_string(vpp_size)); + } + return specs; + } + if (partition.empty()) { + return std::monostate{}; + } + auto counts = ParsePipelineLayerPartition(partition); + if (counts.size() != static_cast(pp_size) * vpp_size) { + throw std::invalid_argument("pipeline_layer_partition needs PP * vPP entries in global chunk order"); + } + return counts; +} + +std::shared_ptr ResolvePipelineLayout(LayerIndex num_layers, int pp_size, int vpp_size, + const PipelineLayoutRequest &request) { + if (pp_size <= 0 || vpp_size <= 0 + || static_cast(pp_size) * vpp_size > std::numeric_limits::max()) { + throw std::invalid_argument("PP/vPP must be positive and their product must fit int"); + } + if (const auto *specs = std::get_if>(&request)) { + auto layout = PipelineLayout::BuildChunkLayout(num_layers, pp_size, *specs); + if (layout.GetMaxLocalChunks() != vpp_size) { + throw std::invalid_argument("explicit layout does not match virtual_pipeline_parallel"); + } + return std::make_shared(std::move(layout)); + } + if (const auto *counts = std::get_if>(&request)) { + if (counts->empty()) { + throw std::invalid_argument("custom partition must not be empty"); + } + return BuildPipelineLayoutFromCLI(num_layers, pp_size, *counts, vpp_size); + } + return BuildPipelineLayoutFromCLI(num_layers, pp_size, {}, vpp_size); +} + +std::string FormatPipelineLayout(const PipelineLayout &layout) { + std::ostringstream output; + output << (layout.IsCustom() ? "custom" : "uniform") << " pipeline layout:"; + + for (int stage_id = 0; stage_id < layout.GetNumStages(); ++stage_id) { + const auto &stage = layout.GetStage(stage_id); + output << "\n Stage " << stage_id << ":"; + for (const auto &chunk : stage.chunks) { + output << " chunk " << chunk.global_chunk_id << " (local " << chunk.local_chunk_idx << ") layers [" + << chunk.layer_range.begin << "," << chunk.layer_range.end << ")"; + } + if (stage.has_embedding) { + output << " embedding"; + } + if (stage.has_final_norm) { + output << " final_norm"; + } + if (stage.has_lm_head) { + output << " lm_head"; + } + } + return output.str(); +} + +} // namespace infini_train::examples diff --git a/example/common/parser.h b/example/common/parser.h new file mode 100644 index 000000000..af2fa8f6b --- /dev/null +++ b/example/common/parser.h @@ -0,0 +1,34 @@ +#pragma once + +#include +#include +#include +#include +#include +#include + +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" + +namespace infini_train::examples { + +using PipelineLayoutRequest + = std::variant, std::vector>; + +PipelineLayoutRequest ParsePipelineLayoutRequest(std::string_view partition, std::string_view chunks, int pp_size, + int vpp_size); + +bool IsExplicitChunkRequest(const PipelineLayoutRequest &request); + +std::shared_ptr ResolvePipelineLayout(nn::parallel::LayerIndex num_layers, + int pp_size, int vpp_size, + const PipelineLayoutRequest &request); + +std::vector ParsePipelineLayerPartition(std::string_view value); + +std::shared_ptr +BuildPipelineLayoutFromCLI(nn::parallel::LayerIndex num_layers, int num_stages, + std::span counts, int chunks_per_stage = 1); + +std::string FormatPipelineLayout(const nn::parallel::PipelineLayout &layout); + +} // namespace infini_train::examples diff --git a/example/gpt2/checkpoint_loader.cc b/example/gpt2/checkpoint_loader.cc index 95e54730b..d3c5cc1bb 100644 --- a/example/gpt2/checkpoint_loader.cc +++ b/example/gpt2/checkpoint_loader.cc @@ -5,9 +5,12 @@ #include #include #include +#include #include +#include #include #include +#include #include #include "glog/logging.h" @@ -22,6 +25,7 @@ #include "infini_train/include/nn/parallel/tensor_parallel.h" #include "infini_train/include/tensor.h" +#include "example/common/parser.h" #include "example/common/utils.h" #include "example/gpt2/config.h" @@ -56,8 +60,13 @@ std::tuple DetermineAndCheckVersion(const std:: } // namespace namespace gpt2 { +namespace { +using LayerIndex = nn::parallel::LayerIndex; +using PipelineLayout = nn::parallel::PipelineLayout; -std::shared_ptr LoadFromLLMC(const std::string &filepath) { +std::shared_ptr +LoadFromLLMCImpl(const std::string &filepath, + std::optional layout_request) { if (!std::filesystem::exists(filepath)) { LOG(FATAL) << "File not found: " << filepath; } @@ -80,16 +89,29 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) const auto padded_vocab_size = BytesToType(header, 28); // NOTE(zbl): vocab_size needs to be padded to multiple of TP size const auto model_vocab_size = tp_size > 1 ? padded_vocab_size : vocab_size; - + // ========== pp_size:num_stages; vpp_size: num_chunks_per_stage ========== + int pp_size = nn::parallel::global::GetPipelineParallelSize(); + int vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); + auto pp_rank = nn::parallel::pp_rank; nn::TransformerConfig gpt2_config = gpt2::GPT2Config(); gpt2_config.block_size = block_size; gpt2_config.vocab_size = model_vocab_size; gpt2_config.original_vocab_size = vocab_size; gpt2_config.n_layer = n_layer; gpt2_config.n_head = n_head; + gpt2_config.n_kv_head = n_head; gpt2_config.n_embd = n_embd; + // Cross-rank tying is not implemented. Keep every PP>1 layout on the + // existing split-weight semantics, even when both endpoints share a rank. + gpt2_config.tie_weights = gpt2_config.tie_weights && pp_size == 1; gpt2::SanitizeGPT2Config(gpt2_config); - auto local_gpt2 = std::make_shared(gpt2_config); + std::shared_ptr pipeline_layout; + if (layout_request.has_value()) { + pipeline_layout = infini_train::examples::ResolvePipelineLayout(static_cast(n_layer), pp_size, + vpp_size, *layout_request); + } + auto local_gpt2 = pipeline_layout ? std::make_shared(gpt2_config, pipeline_layout, pp_rank) + : std::make_shared(gpt2_config); LOG(INFO) << "magic: " << magic << " version: " << version << " block_size: " << block_size << " vocab_size: " << vocab_size << " n_layer: " << n_layer << " n_head: " << n_head @@ -99,16 +121,42 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) CHECK_EQ(n_embd % n_head, 0) << "n_embd must be divisible by n_head."; CHECK_EQ(n_head % tp_size, 0) << "n_head must be divisible by TP world size."; - // ========== pp_size:num_stages; vpp_size: num_chunks_per_stage ========== - int pp_size = nn::parallel::global::GetPipelineParallelSize(); - int vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); - auto pp_rank = nn::parallel::pp_rank; - auto [is_first_stage, is_last_stage, layer_ranges_per_chunk] - = nn::parallel::PipelineParallel::GetStageInfo(n_layer, pp_size, pp_rank, vpp_size); - // ========== layer to chunk ========== + bool has_embedding = false; + bool has_final_norm = false; + bool has_lm_head = false; + std::vector owned_layers(n_layer, false); - for (const auto &[start, end] : layer_ranges_per_chunk) { - for (int i = start; i < end; ++i) { owned_layers[i] = true; } + std::vector local_indices(n_layer, -1); + + if (pipeline_layout) { + const auto &stage = pipeline_layout->GetStage(pp_rank); + has_embedding = stage.has_embedding; + has_final_norm = stage.has_final_norm; + has_lm_head = stage.has_lm_head; + + for (const auto &chunk : stage.chunks) { + for (LayerIndex layer = chunk.layer_range.begin; layer < chunk.layer_range.end; ++layer) { + owned_layers.at(layer) = true; + local_indices.at(layer) = pipeline_layout->GetLocalLayerIndex(pp_rank, layer); + } + } + } else { + const auto legacy = nn::parallel::PipelineParallel::GetStageInfo(n_layer, pp_size, pp_rank, vpp_size); + + has_embedding = legacy.is_first_stage; + has_final_norm = legacy.is_last_stage; + has_lm_head = legacy.is_last_stage; + + for (const auto &[begin, end] : legacy.layer_ranges_per_chunk) { + for (int layer = begin; layer < end; ++layer) { owned_layers.at(layer) = true; } + } + + LayerIndex local = 0; + for (std::size_t layer = 0; layer < owned_layers.size(); ++layer) { + if (owned_layers[layer]) { + local_indices[layer] = local++; + } + } } auto tp_rank = nn::parallel::tp_rank; @@ -130,28 +178,30 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) // transformer.wte.weight (also transformer.lm_head.weight) // full: (model_vocab_size, n_embd) // local: (vocab_size_per_partition, n_embd) - if (is_first_stage) { + const auto wte_offset = ifs.tellg(); + if (has_embedding) { auto &transformer_wte_weight = state_dict[std::format("{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerFirstStage::kWTELayerName, nn::parallel::VocabParallelEmbedding::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(transformer_wte_weight->DataPtr()), model_vocab_size, n_embd, v_start, vpp); - } else if (pp_size > 1 && is_last_stage) { + } + if (has_lm_head) { + ifs.seekg(wte_offset); auto &lm_head_weight = state_dict[std::format("{}.{}", nn::TransformerLastStage::kLMHeadLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(lm_head_weight->DataPtr()), model_vocab_size, n_embd, v_start, vpp); - } else { - size_t wte_bytes = model_vocab_size * n_embd * sizeof(float); - ifs.seekg(wte_bytes, std::ios::cur); } + ifs.seekg(wte_offset + + static_cast(static_cast(model_vocab_size) * n_embd * sizeof(float))); if (tp_size == 1) { // Skip padded vocab part when TP is not enabled ifs.ignore((padded_vocab_size - model_vocab_size) * n_embd * sizeof(float)); } - if (is_first_stage) { + if (has_embedding) { // transformer.wpe.weight auto &transformer_wpe_weight = state_dict[std::format("{}.{}.{}", nn::TransformerModel::kTransformerModelName, @@ -163,15 +213,14 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(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)]; 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 +228,13 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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,9 +242,9 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, @@ -229,7 +277,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,9 +284,9 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, @@ -271,7 +318,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 +325,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 idx = 0; idx < n_layer; ++idx) { - if (owned_layers[idx]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(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)]; 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 +341,14 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(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)]; 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 +356,14 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(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)]; 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 +371,13 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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 +385,14 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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 +400,14 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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 +415,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 idx = 0; idx < n_layer; ++idx) { - if (owned_layers[idx]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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,22 +431,21 @@ 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]) { + if (owned_layers.at(idx)) { + const LayerIndex local_layer_index = local_indices.at(idx); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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); } } - if (is_last_stage) { + if (has_final_norm) { // transformer.ln_f.weight auto &transformer_ln_f_weight = state_dict[std::format("{}.{}.{}", nn::TransformerModel::kTransformerModelName, @@ -426,4 +464,13 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) return local_gpt2; } +} // namespace + +std::shared_ptr LoadFromLLMC(const std::string &filepath) { + return LoadFromLLMCImpl(filepath, std::nullopt); +} +std::shared_ptr LoadFromLLMC(const std::string &filepath, + const infini_train::examples::PipelineLayoutRequest &request) { + return LoadFromLLMCImpl(filepath, request); +} } // namespace gpt2 diff --git a/example/gpt2/checkpoint_loader.h b/example/gpt2/checkpoint_loader.h index e80c356e3..d27807a61 100644 --- a/example/gpt2/checkpoint_loader.h +++ b/example/gpt2/checkpoint_loader.h @@ -1,5 +1,7 @@ #pragma once +#include "example/common/parser.h" + #include #include @@ -9,4 +11,7 @@ class TransformerModel; namespace gpt2 { std::shared_ptr LoadFromLLMC(const std::string &filepath); +std::shared_ptr +LoadFromLLMC(const std::string &filepath, const infini_train::examples::PipelineLayoutRequest &layout_request); + } // namespace gpt2 diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 8e5d92c02..02f1064b4 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -1,12 +1,18 @@ +#include "example/common/parser.h" #include +#include #include +#include #include #include +#include #include #include +#include #include #include #include +#include #include "gflags/gflags.h" #include "glog/logging.h" @@ -136,21 +142,22 @@ DEFINE_validator(zero_stage, [](const char *, int32_t value) { return value >= 0 DEFINE_validator(lr_decay_style, [](const char *, const std::string &value) { return kSupportedLRDecayStyles.contains(value); }); -void Train(const nn::parallel::Rank &rank) { +DEFINE_string(pipeline_layer_partition, "", "Layer counts in global chunk order; requires PP * vPP entries."); +DEFINE_string(pipeline_chunk_layout, "", "Ordered stage:layer_count entries with explicit chunk owners."); + +void Train(const nn::parallel::Rank &rank, examples::PipelineLayoutRequest layout_request) try { using namespace nn::parallel; { if (rank.IsLastRank()) { if (!FLAGS_save.empty() && FLAGS_save_interval == 0) { LOG(FATAL) << "Invalid configuration: --save is set ('" << FLAGS_save - << "'), but --save_interval is 0. " - << "They must be set together."; + << "'), but --save_interval is 0. " << "They must be set together."; } if (FLAGS_save.empty() && FLAGS_save_interval > 0) { LOG(FATAL) << "Invalid configuration: --save_interval is set to " << FLAGS_save_interval - << ", but --save is empty. " - << "They must be set together."; + << ", but --save is empty. " << "They must be set together."; } } } @@ -162,7 +169,7 @@ void Train(const nn::parallel::Rank &rank) { int tp_world_size = global::GetTensorParallelSize(); int sp_world_size = global::GetSequenceParallelEnabled() ? tp_world_size : 1; int pp_world_size = global::GetPipelineParallelSize(); - + const bool explicit_chunk_mode = examples::IsExplicitChunkRequest(layout_request); if (FLAGS_sequence_parallel) { CHECK_EQ(FLAGS_sequence_length % tp_world_size, 0) << "sequence_length must be divisible by tp_world_size when SP is enabled (pad later if needed)."; @@ -224,20 +231,33 @@ void Train(const nn::parallel::Rank &rank) { std::shared_ptr model = nullptr; if (!FLAGS_llmc_filepath.empty()) { - model = gpt2::LoadFromLLMC(FLAGS_llmc_filepath); - } else if (kModelToConfigs.count(FLAGS_model)) { - model_config = kModelToConfigs.at(FLAGS_model); + model = gpt2::LoadFromLLMC(FLAGS_llmc_filepath, layout_request); + } else if (FLAGS_model == "gpt2" || kModelToConfigs.count(FLAGS_model)) { + if (kModelToConfigs.count(FLAGS_model)) { + model_config = kModelToConfigs.at(FLAGS_model); + } + model_config.n_kv_head = model_config.n_head; + // Cross-rank tying is not implemented. Keep every PP>1 layout on the + // existing split-weight semantics, even when both endpoints share a rank. + model_config.tie_weights = model_config.tie_weights && pp_world_size == 1; gpt2::SanitizeGPT2Config(model_config); - model = std::make_shared(model_config); + const auto layout = examples::ResolvePipelineLayout(model_config.n_layer, pp_world_size, + FLAGS_virtual_pipeline_parallel, layout_request); + model = layout ? std::make_shared(model_config, layout, pp_rank) + : std::make_shared(model_config); } + CHECK(model) << "Unable to create GPT-2 model."; + const auto gpt2_model = std::dynamic_pointer_cast(model); + CHECK(gpt2_model) << "GPT2 example expects GPT2 model."; + model_config = gpt2_model->Config(); + const auto pipeline_layout = gpt2_model->GetPipelineLayout(); + if (rank.GlobalRank() == 0 && pipeline_layout) { + LOG(INFO) << examples::FormatPipelineLayout(*pipeline_layout); + } model->To(device); - utils::PrecisionChecker::BuildNameMap(model.get()); - - // Get chunk size before wrapping with LoRA (needed for PipelineParallel) - auto gpt2_model = std::dynamic_pointer_cast(model); - CHECK(gpt2_model) << "GPT2 example expects GPT2 model."; + utils::PrecisionChecker::BuildNameMap(model.get(), pipeline_layout, pp_rank); // Apply LoRA using GetLoRAModel (in-place injection) bool lora_enabled = FLAGS_lora_rank > 0; @@ -287,8 +307,13 @@ void Train(const nn::parallel::Rank &rank) { auto shapes = std::vector>{ {FLAGS_batch_size, FLAGS_sequence_length / sp_world_size, model_config.n_embd}}; - model = std::make_shared(model, pp_world_size, num_micro_batches, shapes, - pp_rank, device, model_config.GetChunkSize()); + if (pipeline_layout) { + model = std::make_shared( + model, pp_world_size, num_micro_batches, shapes, pp_rank, device, pipeline_layout, explicit_chunk_mode); + } else { + model = std::make_shared(model, pp_world_size, num_micro_batches, shapes, + pp_rank, device, model_config.GetChunkSize()); + } if (ddp_world_size > 1) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; auto *mutable_chunks = dynamic_cast(model.get())->mutable_chunks(); @@ -365,13 +390,12 @@ void Train(const nn::parallel::Rank &rank) { auto train_iter = train_loader.begin(); std::shared_ptr loss_fn = (tp_world_size > 1) ? std::static_pointer_cast( - std::make_shared(model_config.original_vocab_size)) + std::make_shared(model_config.original_vocab_size)) : std::static_pointer_cast(std::make_shared()); loss_fn->To(device); LOG(INFO) << "Rank " << rank.GlobalRank() << ": start training"; auto impl = core::GetDeviceGuardImpl(device.type()); - int start_step = 0; TrainerState state; const auto resume_result = ResumeFromCheckpoint({.resume_root = FLAGS_load, @@ -535,7 +559,10 @@ void Train(const nn::parallel::Rank &rank) { const double duration_us = std::chrono::duration(iter_end - iter_start).count(); const double tps = FLAGS_total_batch_size / (duration_us / 1e6); - if (rank.IsLastRank()) { + // PP loss is local: only the logical output stage has a valid value. + // Select one DP/TP replica for logging; the DP AllReduce above remains collective. + const int output_stage = pipeline_layout ? pipeline_layout->GetOutputStage() : pp_world_size - 1; + if (pp_rank == output_stage && ddp_rank == ddp_world_size - 1 && tp_rank == tp_world_size - 1) { size_t used_mb = 0, reserved_mb = 0; std::tie(used_mb, reserved_mb) = impl->GetMemPoolPeakMB(device); LOG(ERROR) << std::format("step {:4d}/{} | train loss {:.6f} | lr {:.2e} | ({:.2f} ms | {:.0f} tok/s | " @@ -544,7 +571,7 @@ void Train(const nn::parallel::Rank &rank) { used_mb, reserved_mb, ddp_world_size, tp_world_size, sp_world_size, pp_world_size); - if ((step + 1) % FLAGS_freq_generate_txt == 0) { + if (FLAGS_freq_generate_txt > 0 && (step + 1) % FLAGS_freq_generate_txt == 0) { if (tokenizer) { // FIXME(jym): to support PP CHECK_EQ(pp_world_size, 1); @@ -575,12 +602,39 @@ void Train(const nn::parallel::Rank &rank) { Profiler::Instance().Report("gpt2.report", Profiler::SortBy::DeviceTimePercentage); Profiler::Instance().PrintRecords("gpt2.records.log"); #endif +} catch (const std::exception &error) { + LOG(FATAL) << "Rank " << rank.GlobalRank() << ": --pipeline_layer_partition=\"" << FLAGS_pipeline_layer_partition + << "\", --pipeline_chunk_layout=\"" << FLAGS_pipeline_chunk_layout << "\": " << error.what(); } int main(int argc, char *argv[]) { gflags::ParseCommandLineFlags(&argc, &argv, true); google::InitGoogleLogging(argv[0]); + examples::PipelineLayoutRequest layout_request; + try { + const auto limit = static_cast(std::numeric_limits::max()); + if (FLAGS_pipeline_parallel == 0 || FLAGS_virtual_pipeline_parallel == 0 || FLAGS_pipeline_parallel > limit + || FLAGS_virtual_pipeline_parallel > limit + || FLAGS_pipeline_parallel > limit / FLAGS_virtual_pipeline_parallel) { + throw std::invalid_argument("PP/vPP must be positive and their product must fit int"); + } + layout_request + = examples::ParsePipelineLayoutRequest(FLAGS_pipeline_layer_partition, FLAGS_pipeline_chunk_layout, + FLAGS_pipeline_parallel, FLAGS_virtual_pipeline_parallel); + if (examples::IsExplicitChunkRequest(layout_request) + && (FLAGS_pipeline_parallel < 2 || FLAGS_sequence_parallel || FLAGS_freq_generate_txt != 0 + || FLAGS_val_loss_every != 0 || FLAGS_sample_every != 0)) { + throw std::invalid_argument("explicit chunks require PP>=2, SP off, and generation/validation disabled"); + } + } catch (const std::exception &error) { + LOG(ERROR) << "Invalid --pipeline_layer_partition=" << FLAGS_pipeline_layer_partition + << ", --pipeline_chunk_layout=" << FLAGS_pipeline_chunk_layout << ", PP=" << FLAGS_pipeline_parallel + << ", vPP=" << FLAGS_virtual_pipeline_parallel << ": " << error.what(); + google::ShutdownGoogleLogging(); + return EXIT_FAILURE; + } + auto precision_config = utils::PrecisionCheckConfig::Parse(FLAGS_precision_check); nn::parallel::global::InitAllEnv(FLAGS_nthread_per_process, FLAGS_tensor_parallel, FLAGS_sequence_parallel, FLAGS_pipeline_parallel, FLAGS_virtual_pipeline_parallel); @@ -593,14 +647,14 @@ int main(int argc, char *argv[]) { for (int idx = 0; idx < FLAGS_nthread_per_process; ++idx) { nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), idx, nn::parallel::global::GetNprocPerNode(), FLAGS_nthread_per_process); - threads.emplace_back(Train, rank); + threads.emplace_back(Train, rank, layout_request); } for (auto &thread : threads) { thread.join(); } } else { nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), FLAGS_nthread_per_process); - Train(rank); + Train(rank, layout_request); } gflags::ShutDownCommandLineFlags(); diff --git a/example/llama3/checkpoint_loader.cc b/example/llama3/checkpoint_loader.cc index f3590af6e..814e46e08 100644 --- a/example/llama3/checkpoint_loader.cc +++ b/example/llama3/checkpoint_loader.cc @@ -1,4 +1,8 @@ #include "example/llama3/checkpoint_loader.h" +#include "example/common/parser.h" +#include +#include +#include #include #include @@ -39,8 +43,13 @@ constexpr int32_t kLLaMA3FP32Version = 3; } // namespace namespace llama3 { +namespace { +using LayerIndex = nn::parallel::LayerIndex; +using PipelineLayout = nn::parallel::PipelineLayout; -std::shared_ptr LoadFromLLMC(const std::string &filepath) { +std::shared_ptr +LoadFromLLMCImpl(const std::string &filepath, + std::optional layout_request) { if (!std::filesystem::exists(filepath)) { LOG(FATAL) << "File not found: " << filepath; } @@ -71,6 +80,7 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) nn::TransformerConfig llama3_config = llama3::LLaMA3Config(); llama3_config.block_size = block_size; llama3_config.vocab_size = vocab_size; + llama3_config.original_vocab_size = vocab_size; llama3_config.n_layer = n_layer; llama3_config.n_head = n_head; llama3_config.n_kv_head = n_kv_head; @@ -82,18 +92,56 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) llama3_config.norm_eps = norm_eps; llama3_config.max_gen_batch_size = max_gen_bs; llama3::SanitizeLLaMA3Config(llama3_config); - auto llama3 = std::make_shared(llama3_config); - - // ========== pp_size:num_stages; vpp_size: num_chunks_per_stage ========== - int pp_size = nn::parallel::global::GetPipelineParallelSize(); - int vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); - auto pp_rank = nn::parallel::pp_rank; - auto [is_first_stage, is_last_stage, layer_ranges_per_chunk] - = nn::parallel::PipelineParallel::GetStageInfo(n_layer, pp_size, pp_rank, vpp_size); - // ========== layer to chunk ========== + const int pp_size = nn::parallel::global::GetPipelineParallelSize(); + const int vpp_size = nn::parallel::global::GetVirtualPipelineParallelSize(); + const int pp_rank = nn::parallel::pp_rank; + + std::shared_ptr pipeline_layout; + std::shared_ptr llama3; + if (layout_request.has_value()) { + pipeline_layout = infini_train::examples::ResolvePipelineLayout(static_cast(n_layer), pp_size, + vpp_size, *layout_request); + llama3 = pipeline_layout ? std::make_shared(llama3_config, pipeline_layout, pp_rank) + : std::make_shared(llama3_config); + } else { + llama3 = std::make_shared(llama3_config); + } + + bool has_embedding = false; + bool has_final_norm = false; + bool has_lm_head = false; std::vector owned_layers(n_layer, false); - for (const auto &[start, end] : layer_ranges_per_chunk) { - for (int i = start; i < end; ++i) { owned_layers[i] = true; } + std::vector local_indices(n_layer, -1); + std::vector> stage_ranges; + + if (pipeline_layout) { + const auto &stage = pipeline_layout->GetStage(pp_rank); + has_embedding = stage.has_embedding; + has_final_norm = stage.has_final_norm; + has_lm_head = stage.has_lm_head; + for (const auto &chunk : stage.chunks) { + stage_ranges.emplace_back(chunk.layer_range.begin, chunk.layer_range.end); + for (LayerIndex layer = chunk.layer_range.begin; layer < chunk.layer_range.end; ++layer) { + const auto index = static_cast(layer); + owned_layers.at(index) = true; + local_indices.at(index) = pipeline_layout->GetLocalLayerIndex(pp_rank, layer); + } + } + } else { + const auto legacy = nn::parallel::PipelineParallel::GetStageInfo(n_layer, pp_size, pp_rank, vpp_size); + has_embedding = legacy.is_first_stage; + has_final_norm = legacy.is_last_stage; + has_lm_head = legacy.is_last_stage; + for (const auto &[begin, end] : legacy.layer_ranges_per_chunk) { + stage_ranges.emplace_back(begin, end); + for (int layer = begin; layer < end; ++layer) { owned_layers.at(static_cast(layer)) = true; } + } + LayerIndex local_index = 0; + for (std::size_t layer = 0; layer < owned_layers.size(); ++layer) { + if (owned_layers[layer]) { + local_indices[layer] = local_index++; + } + } } const int tp_size = nn::parallel::global::GetTensorParallelSize(); @@ -122,9 +170,8 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) LOG(INFO) << " version_minor = " << version_minor; LOG(INFO) << "Pipeline Parallel Chunks:"; - for (size_t i = 0; i < layer_ranges_per_chunk.size(); ++i) { - LOG(INFO) << " Chunk " << i << ": layers " << layer_ranges_per_chunk[i].first << " to " - << layer_ranges_per_chunk[i].second; + for (size_t i = 0; i < stage_ranges.size(); ++i) { + LOG(INFO) << " Chunk " << i << ": layers " << stage_ranges[i].first << " to " << stage_ranges[i].second; } } @@ -167,8 +214,9 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) auto state_dict = llama3->StateDict(); // ========== Read Sharded Params ========== - // transformer.wte.weight : (vocab_size, n_embd) -> local tp_rank: rows of [v_start : v_start+vpp) - if (is_first_stage) { + // transformer.wte.weight : (vocab_size, n_embd) -> local tp_rank: rows of + // [v_start : v_start+vpp) + if (has_embedding) { auto &wte = state_dict[std::format("{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerFirstStage::kWTELayerName, nn::parallel::VocabParallelEmbedding::kParamWeightName)]; @@ -181,25 +229,27 @@ 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]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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); } } - // 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; + // 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 for (int i = 0; i < static_cast(n_layer); ++i) { - if (owned_layers[i]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, @@ -213,33 +263,37 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) ReadMatrixRowShardFloat(ifs, /*dst=*/dst + (0 * attn_cols), /*rows=*/attn_rows_all, /*cols=*/attn_cols, - /*row_start=*/tp_rank * q_local_rows, /*row_cnt=*/q_local_rows); + /*row_start=*/tp_rank * q_local_rows, + /*row_cnt=*/q_local_rows); // K block -> [q_local_rows : q_local_rows + kv_local_rows) ifs.seekg(base_pos); ReadMatrixRowShardFloat(ifs, /*dst=*/dst + (q_local_rows * attn_cols), /*rows=*/attn_rows_all, /*cols=*/attn_cols, - /*row_start=*/q_out_rows + tp_rank * kv_local_rows, /*row_cnt=*/kv_local_rows); + /*row_start=*/q_out_rows + tp_rank * kv_local_rows, + /*row_cnt=*/kv_local_rows); - // V block -> [q_local_rows + kv_local_rows : q_local_rows + 2*kv_local_rows) + // V block -> [q_local_rows + kv_local_rows : q_local_rows + + // 2*kv_local_rows) ifs.seekg(base_pos); ReadMatrixRowShardFloat(ifs, /*dst=*/dst + ((q_local_rows + kv_local_rows) * attn_cols), /*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); } } - // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; + // transformer.h.{i}.attn.c_proj.weight : RowParallelLinear, but actually + // applies on "columns" for (int i = 0; i < static_cast(n_layer); ++i) { - if (owned_layers[i]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format( "{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName, @@ -247,7 +301,6 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) 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,24 +308,26 @@ 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]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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); } } - // transformer.h.{i}.mlp.c_fc.weight : ColumnParallelLinear, but actually applies on "rows" - local_layer_index = 0; + // transformer.h.{i}.mlp.c_fc.weight : ColumnParallelLinear, but actually + // applies on "rows" for (int i = 0; i < static_cast(n_layer); ++i) { - if (owned_layers[i]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, @@ -280,17 +335,18 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) 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); } } - // transformer.h.{i}.mlp.c_fc2.weight : ColumnParallelLinear, but actually applies on "rows" - local_layer_index = 0; + // transformer.h.{i}.mlp.c_fc2.weight : ColumnParallelLinear, but actually + // applies on "rows" for (int i = 0; i < static_cast(n_layer); ++i) { - if (owned_layers[i]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFc2LayerName, @@ -298,25 +354,26 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) 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); } } - // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually applies on "columns" - local_layer_index = 0; + // transformer.h.{i}.mlp.c_proj.weight : RowParallelLinear, but actually + // applies on "columns" for (int i = 0; i < static_cast(n_layer); ++i) { - if (owned_layers[i]) { + if (owned_layers.at(static_cast(i))) { + const LayerIndex local_layer_index = local_indices.at(static_cast(i)); + CHECK_GE(local_layer_index, 0); auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), 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; + /*col_start=*/tp_rank * in_fc_pp, + /*col_cnt=*/in_fc_pp); } else { size_t c_proj_bytes = static_cast(n_embd) * ffn_hidden * sizeof(float); ifs.seekg(c_proj_bytes, std::ios::cur); @@ -324,9 +381,11 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) } // transformer.ln_f.weight : Full version nn::RMSNorm - // lm_head.weight : (vocab_size, n_embd) -> ColumnParallelLinear, but actually applies on "rows" + // lm_head.weight : (vocab_size, n_embd) -> ColumnParallelLinear, but actually + // applies on "rows" + CHECK_EQ(has_final_norm, has_lm_head) << "current combined output module requires final norm and LM head together"; { - if (is_last_stage) { + if (has_final_norm && has_lm_head) { auto &ln_f = state_dict[std::format("{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerLastStage::kLnFLayerName, nn::RMSNorm::kParamWeightName)]; @@ -345,4 +404,15 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) return llama3; } +} // namespace + +std::shared_ptr LoadFromLLMC(const std::string &filepath) { + return LoadFromLLMCImpl(filepath, std::nullopt); +} + +std::shared_ptr +LoadFromLLMC(const std::string &filepath, const infini_train::examples::PipelineLayoutRequest &layout_request) { + return LoadFromLLMCImpl(filepath, layout_request); +} + } // namespace llama3 diff --git a/example/llama3/checkpoint_loader.h b/example/llama3/checkpoint_loader.h index d4aea3d04..03f984428 100644 --- a/example/llama3/checkpoint_loader.h +++ b/example/llama3/checkpoint_loader.h @@ -1,5 +1,7 @@ #pragma once +#include "example/common/parser.h" + #include #include @@ -9,4 +11,7 @@ class TransformerModel; namespace llama3 { std::shared_ptr LoadFromLLMC(const std::string &filepath); +std::shared_ptr +LoadFromLLMC(const std::string &filepath, const infini_train::examples::PipelineLayoutRequest &layout_request); + } // namespace llama3 diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 19620e993..ccdf6a791 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -1,10 +1,16 @@ +#include "example/common/parser.h" +#include #include +#include #include #include +#include #include #include +#include #include #include +#include #include "gflags/gflags.h" #include "glog/logging.h" @@ -125,21 +131,22 @@ DEFINE_validator(zero_stage, [](const char *, int32_t value) { return value >= 0 DEFINE_validator(lr_decay_style, [](const char *, const std::string &value) { return kSupportedLRDecayStyles.contains(value); }); -void Train(const nn::parallel::Rank &rank) { +DEFINE_string(pipeline_layer_partition, "", "Layer counts in global chunk order; requires PP * vPP entries."); +DEFINE_string(pipeline_chunk_layout, "", "Ordered stage:layer_count entries with explicit chunk owners."); + +void Train(const nn::parallel::Rank &rank, examples::PipelineLayoutRequest layout_request) try { using namespace nn::parallel; { if (rank.IsLastRank()) { if (!FLAGS_save.empty() && FLAGS_save_interval == 0) { LOG(FATAL) << "Invalid configuration: --save is set ('" << FLAGS_save - << "'), but --save_interval is 0. " - << "They must be set together."; + << "'), but --save_interval is 0. " << "They must be set together."; } if (FLAGS_save.empty() && FLAGS_save_interval > 0) { LOG(FATAL) << "Invalid configuration: --save_interval is set to " << FLAGS_save_interval - << ", but --save is empty. " - << "They must be set together."; + << ", but --save is empty. " << "They must be set together."; } } } @@ -151,7 +158,7 @@ void Train(const nn::parallel::Rank &rank) { int tp_world_size = global::GetTensorParallelSize(); int sp_world_size = global::GetSequenceParallelEnabled() ? tp_world_size : 1; int pp_world_size = global::GetPipelineParallelSize(); - + const bool explicit_chunk_mode = examples::IsExplicitChunkRequest(layout_request); if (FLAGS_sequence_parallel) { CHECK_EQ(FLAGS_sequence_length % tp_world_size, 0) << "sequence_length must be divisible by tp_world_size when SP is enabled (pad later if needed)."; @@ -212,15 +219,27 @@ void Train(const nn::parallel::Rank &rank) { nn::TransformerConfig model_config = llama3::LLaMA3Config(); std::shared_ptr model = nullptr; if (!FLAGS_llmc_filepath.empty()) { - model = llama3::LoadFromLLMC(FLAGS_llmc_filepath); + model = llama3::LoadFromLLMC(FLAGS_llmc_filepath, layout_request); } else { llama3::SanitizeLLaMA3Config(model_config); - model = std::make_shared(model_config); + const auto layout = examples::ResolvePipelineLayout(model_config.n_layer, pp_world_size, + FLAGS_virtual_pipeline_parallel, layout_request); + model = layout ? std::make_shared(model_config, layout, pp_rank) + : std::make_shared(model_config); + } + + const auto llama3_model = std::dynamic_pointer_cast(model); + CHECK(llama3_model) << "LLaMA 3 example expects TransformerModel."; + model_config = llama3_model->Config(); + + const auto pipeline_layout = llama3_model->GetPipelineLayout(); + if (rank.GlobalRank() == 0 && pipeline_layout) { + LOG(INFO) << examples::FormatPipelineLayout(*pipeline_layout); } model->To(device); - utils::PrecisionChecker::BuildNameMap(model.get()); + utils::PrecisionChecker::BuildNameMap(model.get(), pipeline_layout, pp_rank); // Apply LoRA using GetLoRAModel (in-place injection) bool lora_enabled = FLAGS_lora_rank > 0; @@ -260,8 +279,13 @@ void Train(const nn::parallel::Rank &rank) { auto shapes = std::vector>{ {FLAGS_batch_size, FLAGS_sequence_length / sp_world_size, model_config.n_embd}}; - model = std::make_shared(model, pp_world_size, num_micro_batches, shapes, - pp_rank, device, model_config.GetChunkSize()); + if (pipeline_layout) { + model = std::make_shared( + model, pp_world_size, num_micro_batches, shapes, pp_rank, device, pipeline_layout, explicit_chunk_mode); + } else { + model = std::make_shared(model, pp_world_size, num_micro_batches, shapes, + pp_rank, device, model_config.GetChunkSize()); + } if (ddp_world_size > 1) { auto ddp_config = DistributedDataParallelConfig{.zero_stage = FLAGS_zero_stage}; auto *mutable_chunks = dynamic_cast(model.get())->mutable_chunks(); @@ -352,7 +376,6 @@ void Train(const nn::parallel::Rank &rank) { LOG(INFO) << "Rank " << rank.GlobalRank() << ": start training"; auto impl = core::GetDeviceGuardImpl(device.type()); - int start_step = 0; TrainerState state; const auto resume_result = ResumeFromCheckpoint({.resume_root = FLAGS_load, @@ -514,7 +537,10 @@ void Train(const nn::parallel::Rank &rank) { const double duration_us = std::chrono::duration(iter_end - iter_start).count(); const double tps = FLAGS_total_batch_size / (duration_us / 1e6); - if (rank.IsLastRank()) { + // PP loss is local: only the logical output stage has a valid value. + // Select one DP/TP replica for logging; the DP AllReduce above remains collective. + const int output_stage = pipeline_layout ? pipeline_layout->GetOutputStage() : pp_world_size - 1; + if (pp_rank == output_stage && ddp_rank == ddp_world_size - 1 && tp_rank == tp_world_size - 1) { size_t used_mb = 0, reserved_mb = 0; std::tie(used_mb, reserved_mb) = impl->GetMemPoolPeakMB(device); LOG(ERROR) << std::format("step {:4d}/{} | train loss {:.6f} | lr {:.2e} | ({:.2f} ms | {:.0f} tok/s | " @@ -523,7 +549,7 @@ void Train(const nn::parallel::Rank &rank) { used_mb, reserved_mb, ddp_world_size, tp_world_size, sp_world_size, pp_world_size); - if ((step + 1) % FLAGS_freq_generate_txt == 0) { + if (FLAGS_freq_generate_txt > 0 && (step + 1) % FLAGS_freq_generate_txt == 0) { // FIXME(jym): to support PP if (tokenizer) { CHECK_EQ(pp_world_size, 1); @@ -554,12 +580,39 @@ void Train(const nn::parallel::Rank &rank) { Profiler::Instance().Report("llama3.report", Profiler::SortBy::DeviceTimePercentage); Profiler::Instance().PrintRecords("llama3.records.log"); #endif +} catch (const std::exception &error) { + LOG(FATAL) << "Rank " << rank.GlobalRank() << ": --pipeline_layer_partition=\"" << FLAGS_pipeline_layer_partition + << "\", --pipeline_chunk_layout=\"" << FLAGS_pipeline_chunk_layout << "\": " << error.what(); } int main(int argc, char *argv[]) { gflags::ParseCommandLineFlags(&argc, &argv, true); google::InitGoogleLogging(argv[0]); + examples::PipelineLayoutRequest layout_request; + try { + const auto limit = static_cast(std::numeric_limits::max()); + if (FLAGS_pipeline_parallel == 0 || FLAGS_virtual_pipeline_parallel == 0 || FLAGS_pipeline_parallel > limit + || FLAGS_virtual_pipeline_parallel > limit + || FLAGS_pipeline_parallel > limit / FLAGS_virtual_pipeline_parallel) { + throw std::invalid_argument("PP/vPP must be positive and their product must fit int"); + } + layout_request + = examples::ParsePipelineLayoutRequest(FLAGS_pipeline_layer_partition, FLAGS_pipeline_chunk_layout, + FLAGS_pipeline_parallel, FLAGS_virtual_pipeline_parallel); + if (examples::IsExplicitChunkRequest(layout_request) + && (FLAGS_pipeline_parallel < 2 || FLAGS_sequence_parallel || FLAGS_freq_generate_txt != 0 + || FLAGS_val_loss_every != 0 || FLAGS_sample_every != 0)) { + throw std::invalid_argument("explicit chunks require PP>=2, SP off, and generation/validation disabled"); + } + } catch (const std::exception &error) { + LOG(ERROR) << "Invalid --pipeline_layer_partition=" << FLAGS_pipeline_layer_partition + << ", --pipeline_chunk_layout=" << FLAGS_pipeline_chunk_layout << ", PP=" << FLAGS_pipeline_parallel + << ", vPP=" << FLAGS_virtual_pipeline_parallel << ": " << error.what(); + google::ShutdownGoogleLogging(); + return EXIT_FAILURE; + } + auto precision_config = utils::PrecisionCheckConfig::Parse(FLAGS_precision_check); nn::parallel::global::InitAllEnv(FLAGS_nthread_per_process, FLAGS_tensor_parallel, FLAGS_sequence_parallel, FLAGS_pipeline_parallel, FLAGS_virtual_pipeline_parallel); @@ -572,14 +625,14 @@ int main(int argc, char *argv[]) { for (int idx = 0; idx < FLAGS_nthread_per_process; ++idx) { nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), idx, nn::parallel::global::GetNprocPerNode(), FLAGS_nthread_per_process); - threads.emplace_back(Train, rank); + threads.emplace_back(Train, rank, layout_request); } for (auto &thread : threads) { thread.join(); } } else { nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), FLAGS_nthread_per_process); - Train(rank); + Train(rank, layout_request); } gflags::ShutDownCommandLineFlags(); diff --git a/infini_train/include/nn/modules/transformer/transformer.h b/infini_train/include/nn/modules/transformer/transformer.h index 0471c32fe..90d0bc9be 100644 --- a/infini_train/include/nn/modules/transformer/transformer.h +++ b/infini_train/include/nn/modules/transformer/transformer.h @@ -1,9 +1,11 @@ #pragma once +#include #include #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/modules/transformer/transformer_config.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" #include "infini_train/include/nn/parallel/pp/pipeline_parallel.h" namespace infini_train::nn { @@ -72,15 +74,24 @@ class TransformerModel : public CloneableModule { static constexpr char kTransformerModelName[] = "transformer"; explicit TransformerModel(const TransformerConfig config); + TransformerModel(TransformerConfig config, std::shared_ptr pipeline_layout, + int pipeline_stage_id); std::vector> Forward(const std::vector> &x) override; const TransformerConfig &Config() const { return config_; } + const std::shared_ptr &GetPipelineLayout() const { return pipeline_layout_; } + int GetStageId() const { return pipeline_stage_id_; } + int GetNumChunks() const { return static_cast(stage_info_.layer_ranges_per_chunk.size()); } private: const TransformerConfig config_; + const std::shared_ptr pipeline_layout_; + const int pipeline_stage_id_ = -1; const infini_train::nn::parallel::StageInfo stage_info_; + + void BuildModules(); }; } // namespace infini_train::nn diff --git a/infini_train/include/nn/parallel/pp/pipeline_layout.h b/infini_train/include/nn/parallel/pp/pipeline_layout.h new file mode 100644 index 000000000..f7d6fb730 --- /dev/null +++ b/infini_train/include/nn/parallel/pp/pipeline_layout.h @@ -0,0 +1,62 @@ +#pragma once +#include +#include +#include + +namespace infini_train::nn::parallel { +using LayerIndex = std::int64_t; + +struct LayerRange { + LayerIndex begin = 0; + LayerIndex end = 0; + LayerIndex Size() const; + bool Contains(LayerIndex layer_id) const; +}; +struct PipelineChunkLayout { + int local_chunk_idx = -1; + LayerRange layer_range; + int global_chunk_id = -1; + int stage_id = -1; +}; +struct PipelineChunkSpec { + int stage_id; + LayerIndex layer_count; +}; +struct PipelineStageLayout { + int stage_id = -1; + std::vector chunks; + bool has_embedding = false; + bool has_final_norm = false; + bool has_lm_head = false; +}; +class PipelineLayout { +public: + PipelineLayout() = delete; + static PipelineLayout BuildUniformLayout(LayerIndex num_layers, int num_stages, int chunks_per_stage = 1); + static PipelineLayout BuildCustomLayout(LayerIndex num_layers, int num_stages, std::span counts, + int chunks_per_stage = 1); + static PipelineLayout BuildChunkLayout(LayerIndex num_layers, int num_stages, + std::span chunk_specs); + + const PipelineStageLayout &GetStage(int stage_id) const; + const PipelineChunkLayout &GetChunk(int global_chunk_id) const; + + int GetStageForLayer(LayerIndex global_layer_id) const; + LayerIndex GetLocalLayerIndex(int stage_id, LayerIndex global_layer_id) const; + + int GetNumChunks() const; + int GetInputStage() const; + int GetOutputStage() const; + int GetMaxLocalChunks() const; + LayerIndex GetNumLayers() const; + int GetNumStages() const; + bool IsCustom() const; + +private: + PipelineLayout(LayerIndex num_layers, int num_stages, bool is_custom, std::vector stages); + LayerIndex num_layers_; + int num_stages_; + bool is_custom_; + std::vector stages_; +}; +} // namespace infini_train::nn::parallel diff --git a/infini_train/include/nn/parallel/pp/pipeline_parallel.h b/infini_train/include/nn/parallel/pp/pipeline_parallel.h index 25939bdc2..5952bdf2a 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_parallel.h +++ b/infini_train/include/nn/parallel/pp/pipeline_parallel.h @@ -1,10 +1,12 @@ // pipeline_parallel.h #pragma once +#include #include #include #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" namespace infini_train { class Tensor; @@ -31,6 +33,9 @@ class PipelineParallel : public Module { public: PipelineParallel(const std::shared_ptr module, int num_stages, int num_micro_batches, const std::vector> &recv_shape, int rank, Device device, int vpp); + PipelineParallel(const std::shared_ptr module, int num_stages, int num_micro_batches, + const std::vector> &recv_shape, int rank, Device device, + std::shared_ptr pipeline_layout, bool explicit_chunk_mode = false); float TrainStep(const std::vector> &input, const std::vector> &target, const std::shared_ptr &optimizer, @@ -39,6 +44,7 @@ class PipelineParallel : public Module { static StageInfo GetStageInfo(int total_layers, int pp_size, int pp_rank, int chunks_per_stage = 1); std::vector> *mutable_chunks(); + const std::shared_ptr &GetPipelineLayout() const { return pipeline_layout_; } private: void BuildPipelineStage(const std::vector> &recv_shape, Device device, @@ -50,5 +56,7 @@ class PipelineParallel : public Module { int rank_ = -1; std::shared_ptr schedule_ = nullptr; std::shared_ptr pipeline_stage_ = nullptr; + std::shared_ptr pipeline_layout_ = nullptr; + bool explicit_chunk_mode_ = false; }; } // namespace infini_train::nn::parallel diff --git a/infini_train/include/nn/parallel/pp/pipeline_schedule.h b/infini_train/include/nn/parallel/pp/pipeline_schedule.h index cae190f82..0b065ffea 100644 --- a/infini_train/include/nn/parallel/pp/pipeline_schedule.h +++ b/infini_train/include/nn/parallel/pp/pipeline_schedule.h @@ -16,11 +16,12 @@ class Module; namespace infini_train::nn::parallel { class PipelineStage; +class PipelineLayout; class PipelineSchedule { public: - PipelineSchedule(std::shared_ptr stage, int num_stages, int num_micro_batches) - : stage_(std::move(stage)), num_micro_batches_(num_micro_batches) {} + PipelineSchedule(std::shared_ptr stage, int num_stages, int num_micro_batches, + std::shared_ptr layout = nullptr, bool explicit_chunk_mode = false); virtual ~PipelineSchedule() = default; @@ -35,8 +36,11 @@ class PipelineSchedule { std::vector> SendToNext(const std::vector> &tensors, int peer_rank); protected: + bool has_printed_schedule_ = false; int num_micro_batches_ = -1; std::shared_ptr stage_ = nullptr; + std::shared_ptr layout_; + bool explicit_chunk_mode_ = false; }; class PipelineParallelScheduler { @@ -52,11 +56,14 @@ class PipelineParallelScheduler { bool is_last_chunk; }; - static Task CreateTask(int step, int mb, int global_chunk, int num_stages, int total_chunks, bool is_forward); + static Task CreateTask(int step, int mb, int global_chunk, int num_stages, int total_chunks, bool is_forward, + const PipelineLayout *layout = nullptr); - static std::vector GenerateGPipeSchedule(int n, int num_stages, int vpp_size); + static std::vector GenerateGPipeSchedule(int n, int num_stages, int vpp_size, + const PipelineLayout *layout = nullptr); - static std::vector GenerateInterleaved1F1BSchedule(int n, int num_stages, int vpp_size); + static std::vector GenerateInterleaved1F1BSchedule(int n, int num_stages, int vpp_size, + const PipelineLayout *layout = nullptr); }; } // namespace infini_train::nn::parallel diff --git a/infini_train/include/utils/precision_checker.h b/infini_train/include/utils/precision_checker.h index 87214fb5d..4cbcd8125 100644 --- a/infini_train/include/utils/precision_checker.h +++ b/infini_train/include/utils/precision_checker.h @@ -16,6 +16,9 @@ class Function; namespace nn { class Module; +namespace parallel { +class PipelineLayout; +} } // namespace nn namespace utils { @@ -40,7 +43,8 @@ class PrecisionChecker { // Build name map from root_model without registering hooks // Called by PrecisionCheckEnv::RegisterWithRootModel - static void BuildNameMap(nn::Module *root_model); + static void BuildNameMap(nn::Module *root_model, + std::shared_ptr layout = nullptr, int pp_rank = 0); static void RegisterForFunction(autograd::Function *func, const std::string &name = "", const Config &config = DefaultConfig()); diff --git a/infini_train/src/nn/modules/transformer/transformer.cc b/infini_train/src/nn/modules/transformer/transformer.cc index 99a739d2d..8af0c3dd2 100644 --- a/infini_train/src/nn/modules/transformer/transformer.cc +++ b/infini_train/src/nn/modules/transformer/transformer.cc @@ -1,8 +1,11 @@ #include "infini_train/include/nn/modules/transformer/transformer.h" #include +#include #include -#include +#include +#include +#include #include #include "glog/logging.h" @@ -23,6 +26,46 @@ #include "infini_train/include/tensor.h" namespace infini_train::nn { +namespace { +parallel::StageInfo ConvertLayoutToStageInfo(parallel::LayerIndex num_layers, + const std::shared_ptr &pipeline_layout, + int stage_id) { + if (pipeline_layout == nullptr) { + throw std::invalid_argument("pipeline layout must not be null"); + } + if (pipeline_layout->GetNumLayers() != num_layers) { + throw std::invalid_argument("pipeline_layout layer count does not match TransformerConfig::n_layer"); + } + + const parallel::PipelineStageLayout &stage = pipeline_layout->GetStage(stage_id); + if (stage.has_final_norm != stage.has_lm_head) { + throw std::invalid_argument("final norm and LM head must belong to the same pipeline stage"); + } + + std::vector> layer_ranges_per_chunk; + layer_ranges_per_chunk.reserve(stage.chunks.size()); + for (std::size_t chunk_idx = 0; chunk_idx < stage.chunks.size(); ++chunk_idx) { + const parallel::PipelineChunkLayout &chunk = stage.chunks[chunk_idx]; + if (chunk.stage_id != stage.stage_id || chunk.global_chunk_id < 0) { + throw std::invalid_argument("pipeline chunk identity is inconsistent with its stage"); + } + if (chunk.local_chunk_idx != static_cast(chunk_idx)) { + throw std::invalid_argument("pipeline local chunk indices must be contiguous and start at zero"); + } + if (chunk.layer_range.begin < 0 || chunk.layer_range.end < chunk.layer_range.begin + || chunk.layer_range.end > std::numeric_limits::max()) { + throw std::invalid_argument("pipeline layer range cannot be represented by TransformerChunk"); + } + layer_ranges_per_chunk.emplace_back(static_cast(chunk.layer_range.begin), + static_cast(chunk.layer_range.end)); + } + return parallel::StageInfo{ + .is_first_stage = stage.has_embedding, + .is_last_stage = stage.has_lm_head, + .layer_ranges_per_chunk = std::move(layer_ranges_per_chunk), + }; +} +} // namespace TransformerFirstStage::TransformerFirstStage(const TransformerConfig &config) : CloneableModule(kType), config_(config) { @@ -204,17 +247,30 @@ std::vector> TransformerLastStage::Forward(const std::ve return (*modules_[kLMHeadLayerName])(x1); } -TransformerModel::TransformerModel(const TransformerConfig config) - : CloneableModule(kType), config_(config), +TransformerModel::TransformerModel(TransformerConfig config) + : CloneableModule(kType), config_(config), pipeline_stage_id_(parallel::pp_rank), stage_info_(nn::parallel::PipelineParallel::GetStageInfo( - config_.n_layer, nn::parallel::global::GetPipelineParallelSize(), nn::parallel::pp_rank, + config_.n_layer, nn::parallel::global::GetPipelineParallelSize(), pipeline_stage_id_, nn::parallel::global::GetVirtualPipelineParallelSize())) { + BuildModules(); +} + +TransformerModel::TransformerModel(TransformerConfig config, + std::shared_ptr pipeline_layout, + int pipeline_stage_id) + : CloneableModule(kType), config_(std::move(config)), pipeline_layout_(std::move(pipeline_layout)), + pipeline_stage_id_(pipeline_stage_id), + stage_info_(ConvertLayoutToStageInfo(config_.n_layer, pipeline_layout_, pipeline_stage_id_)) { + BuildModules(); +} + +void TransformerModel::BuildModules() { auto tp_world_size = nn::parallel::global::GetTensorParallelSize(); // NOTE(zbl): VocabParallelEmbedding requires vocab_size % tp_size == 0 // Megatron-LM has an optional argument `--make-vocab-size-divisible-by`, would do padding to vocab // Here we introduce padding by default, might need modify Tokenizer correspondingly later - CHECK_EQ(config.vocab_size % tp_world_size, 0) << "Vocab size should be divisible by TP world size"; + CHECK_EQ(config_.vocab_size % tp_world_size, 0) << "Vocab size should be divisible by TP world size"; std::unordered_map> transformer; if (stage_info_.is_first_stage) { @@ -228,17 +284,11 @@ TransformerModel::TransformerModel(const TransformerConfig config) } { - std::map>> start_layer_to_layer_size_and_chunk; - for (int chunk_idx = 0; chunk_idx < stage_info_.layer_ranges_per_chunk.size(); ++chunk_idx) { - const auto [start_layer, end_layer] = stage_info_.layer_ranges_per_chunk[chunk_idx]; - 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; 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) { + for (const auto &[start_layer, end_layer] : stage_info_.layer_ranges_per_chunk) { + auto chunk = std::make_shared(config_, start_layer, end_layer); + for (int idx = 0; idx < end_layer - start_layer; ++idx) { h.push_back(chunk->mutable_module(TransformerChunk::kHLayerName)->mutable_module(std::to_string(idx))); } modules_[kPPChunkNamePrefix + std::to_string(chunk_idx)] = std::move(chunk); @@ -262,7 +312,7 @@ 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. - if (config_.tie_weights && nn::parallel::global::GetPipelineParallelSize() == 1) { + if (config_.tie_weights && stage_info_.is_first_stage && stage_info_.is_last_stage) { // https://paperswithcode.com/method/weight-tying *mutable_module(kTransformerModelName) ->mutable_module(TransformerFirstStage::kWTELayerName) diff --git a/infini_train/src/nn/parallel/pp/pipeline_layout.cc b/infini_train/src/nn/parallel/pp/pipeline_layout.cc new file mode 100644 index 000000000..ce5501b93 --- /dev/null +++ b/infini_train/src/nn/parallel/pp/pipeline_layout.cc @@ -0,0 +1,260 @@ +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace infini_train::nn::parallel { +namespace { +void ValidateDimensions(LayerIndex num_layers, int num_stages) { + if (num_layers <= 0) { + throw std::invalid_argument("num_layers must be positive"); + } + if (num_stages <= 0) { + throw std::invalid_argument("num_stages must be positive"); + } + if (num_layers < num_stages) { + throw std::invalid_argument("num_layers must be at least num_stages"); + } +} + +std::string FormatLayerCounts(std::span layers_per_stage) { + std::ostringstream output; + output << "["; + for (std::size_t index = 0; index < layers_per_stage.size(); ++index) { + if (index != 0) { + output << ","; + } + output << layers_per_stage[index]; + } + output << "]"; + return output.str(); +} +} // namespace + +LayerIndex LayerRange::Size() const { return end - begin; } + +bool LayerRange::Contains(LayerIndex layer_id) const { return layer_id >= begin && layer_id < end; } + +PipelineLayout::PipelineLayout(LayerIndex num_layers, int num_stages, bool is_custom, + std::vector stages) + : num_layers_(num_layers), num_stages_(num_stages), is_custom_(is_custom), stages_(std::move(stages)) {} + +PipelineLayout PipelineLayout::BuildUniformLayout(LayerIndex num_layers, int num_stages, int chunks_per_stage) { + // Do not call ValidateDimensions(): custom layouts require positive layers and + // at least one layer per stage. The legacy GetStageInfo() adapter also needs + // zero/underfilled default metadata; this does not imply empty-stage training support. + if (num_layers < 0 || num_stages <= 0 || chunks_per_stage <= 0) { + throw std::invalid_argument("uniform layout requires non-negative layers and positive PP/vPP"); + } + + const LayerIndex total_chunks = static_cast(num_stages) * chunks_per_stage; + if (total_chunks > std::numeric_limits::max()) { + throw std::invalid_argument("uniform layout chunk count must fit int"); + } + + const LayerIndex quotient = num_layers / total_chunks; + const LayerIndex remainder = num_layers % total_chunks; + std::vector stages(static_cast(num_stages)); + for (int stage_id = 0; stage_id < num_stages; ++stage_id) { + auto &stage = stages[stage_id]; + stage.stage_id = stage_id; + stage.has_embedding = (stage_id == 0); + stage.has_final_norm = (stage_id == num_stages - 1); + stage.has_lm_head = (stage_id == num_stages - 1); + for (int local_chunk_id = 0; local_chunk_id < chunks_per_stage; ++local_chunk_id) { + const int gid = local_chunk_id * num_stages + stage_id; + const LayerIndex begin + = static_cast(gid) * quotient + + (static_cast(gid) > remainder ? remainder : static_cast(gid)); + const LayerIndex end = begin + quotient + (gid < remainder ? 1 : 0); + + if (begin == end) { + continue; + } + stage.chunks.push_back(PipelineChunkLayout{local_chunk_id, LayerRange{begin, end}, gid, stage_id}); + } + } + + return PipelineLayout(num_layers, num_stages, false, std::move(stages)); +} + +PipelineLayout PipelineLayout::BuildCustomLayout(LayerIndex num_layers, int num_stages, + std::span counts, int chunks_per_stage) { + ValidateDimensions(num_layers, num_stages); + const LayerIndex total_chunks = static_cast(num_stages) * chunks_per_stage; + if (chunks_per_stage <= 0 || total_chunks > std::numeric_limits::max() + || counts.size() != static_cast(total_chunks)) { + throw std::invalid_argument("custom layer count must match num_stages * chunks_per_stage"); + } + LayerIndex actual_sum = 0; + for (LayerIndex count : counts) { + if (count <= 0) { + throw std::invalid_argument("custom layer counts must be positive"); + } + if (count > std::numeric_limits::max() - actual_sum) { + throw std::invalid_argument("custom layer count sum overflow"); + } + actual_sum += count; + } + if (actual_sum != num_layers) { + std::ostringstream msg; + msg << "Invalid pipeline layer counts: received=" << FormatLayerCounts(counts) << ", num_stages=" << num_stages + << ", chunks_per_stage=" << chunks_per_stage << ", expected_sum=" << num_layers << ", actual_sum=" + << actual_sum; + throw std::invalid_argument(msg.str()); + } + + std::vector stages(static_cast(num_stages)); + for (int stage_id = 0; stage_id < num_stages; ++stage_id) { + stages[stage_id].stage_id = stage_id; + stages[stage_id].has_embedding = (stage_id == 0); + stages[stage_id].has_final_norm = (stage_id == num_stages - 1); + stages[stage_id].has_lm_head = (stage_id == num_stages - 1); + } + LayerIndex begin = 0; + for (int gid = 0; gid < static_cast(total_chunks); ++gid) { + const int stage_id = gid % num_stages; + const int local_chunk_id = gid / num_stages; + const LayerIndex end = begin + counts[static_cast(gid)]; + stages[stage_id].chunks.push_back(PipelineChunkLayout{local_chunk_id, LayerRange{begin, end}, gid, stage_id}); + begin = end; + } + return PipelineLayout(num_layers, num_stages, true, std::move(stages)); +} + +PipelineLayout PipelineLayout::BuildChunkLayout(LayerIndex num_layers, int num_stages, + std::span specs) { + ValidateDimensions(num_layers, num_stages); + if (specs.empty() || specs.size() > static_cast(std::numeric_limits::max())) { + throw std::invalid_argument("invalid explicit chunk count"); + } + std::vector stages(static_cast(num_stages)); + for (int id = 0; id < num_stages; ++id) { stages[id].stage_id = id; } + LayerIndex begin = 0; + int global_id = 0; + for (const auto &spec : specs) { + if (spec.stage_id < 0 || spec.stage_id >= num_stages || spec.layer_count <= 0 + || spec.layer_count > num_layers - begin) { + throw std::invalid_argument("invalid pipeline_chunk_layout entry " + std::to_string(global_id) + + ": stage_id=" + std::to_string(spec.stage_id) + ", count=" + + std::to_string(spec.layer_count) + ", stages=" + std::to_string(num_stages) + + ", remaining_layers=" + std::to_string(num_layers - begin)); + } + auto &stage = stages[spec.stage_id]; + const auto end = begin + spec.layer_count; + stage.chunks.push_back(PipelineChunkLayout{static_cast(stage.chunks.size()), LayerRange{begin, end}, + global_id++, spec.stage_id}); + begin = end; + } + if (begin != num_layers) { + throw std::invalid_argument("explicit chunk layer sum=" + std::to_string(begin) + + ", expected=" + std::to_string(num_layers)); + } + for (const auto &stage : stages) { + if (stage.chunks.empty()) { + throw std::invalid_argument("stage " + std::to_string(stage.stage_id) + " owns no chunk"); + } + } + stages[specs.front().stage_id].has_embedding = true; + stages[specs.back().stage_id].has_final_norm = true; + stages[specs.back().stage_id].has_lm_head = true; + return PipelineLayout(num_layers, num_stages, true, std::move(stages)); +} + +const PipelineStageLayout &PipelineLayout::GetStage(int stage_id) const { + if (stage_id < 0 || stage_id >= num_stages_) { + throw std::out_of_range("stage_id is out of range"); + } + return stages_[static_cast(stage_id)]; +} + +const PipelineChunkLayout &PipelineLayout::GetChunk(int global_chunk_id) const { + if (global_chunk_id < 0) { + throw std::out_of_range("global_chunk_id is out of range"); + } + for (const auto &stage : stages_) { + for (const auto &chunk : stage.chunks) { + if (chunk.global_chunk_id == global_chunk_id) { + return chunk; + } + } + } + throw std::out_of_range("global_chunk_id is out of range"); +} + +int PipelineLayout::GetStageForLayer(LayerIndex layer_id) const { + if (layer_id < 0 || layer_id >= num_layers_) { + throw std::out_of_range("layer_id is out of range"); + } + + for (const PipelineStageLayout &stage : stages_) { + for (const auto &chunk : stage.chunks) { + if (chunk.layer_range.Contains(layer_id)) { + return stage.stage_id; + } + } + } + + throw std::out_of_range("layer_id is not assigned to a pipeline stage"); +} + +LayerIndex PipelineLayout::GetLocalLayerIndex(int stage_id, LayerIndex global_layer_id) const { + const PipelineStageLayout &stage = GetStage(stage_id); + + if (global_layer_id < 0 || global_layer_id >= num_layers_) { + throw std::out_of_range("layer_id is out of range"); + } + + LayerIndex offset = 0; + for (const auto &chunk : stage.chunks) { + if (chunk.layer_range.Contains(global_layer_id)) { + return offset + global_layer_id - chunk.layer_range.begin; + } + offset += chunk.layer_range.Size(); + } + throw std::out_of_range("layer_id does not belong to the requested stage"); +} + +int PipelineLayout::GetNumChunks() const { + int total = 0; + for (const auto &stage : stages_) { total += static_cast(stage.chunks.size()); } + return total; +} + +int PipelineLayout::GetMaxLocalChunks() const { + int count = 0; + for (const auto &stage : stages_) { count = std::max(count, static_cast(stage.chunks.size())); } + return count; +} + +int PipelineLayout::GetInputStage() const { + for (const auto &stage : stages_) { + if (stage.has_embedding) { + return stage.stage_id; + } + } + throw std::logic_error("layout has no input stage"); +} + +int PipelineLayout::GetOutputStage() const { + for (const auto &stage : stages_) { + if (stage.has_final_norm && stage.has_lm_head) { + return stage.stage_id; + } + } + throw std::logic_error("layout has no output stage"); +} + +LayerIndex PipelineLayout::GetNumLayers() const { return num_layers_; } + +int PipelineLayout::GetNumStages() const { return num_stages_; } + +bool PipelineLayout::IsCustom() const { return is_custom_; } + +} // namespace infini_train::nn::parallel diff --git a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc index c0369cdeb..3cca59c2c 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_parallel.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_parallel.cc @@ -1,18 +1,58 @@ // pipeline_parallel.cc #include "infini_train/include/nn/parallel/pp/pipeline_parallel.h" +#include #include #include +#include #include +#include #include "infini_train/include/nn/modules/container.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/modules/transformer/transformer.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" #include "infini_train/include/nn/parallel/pp/pipeline_schedule.h" #include "infini_train/include/nn/parallel/pp/pipeline_stage.h" namespace infini_train::nn::parallel { namespace { constexpr char kModuleName[] = "module"; + +std::vector> BuildPipelineChunksFromLayout(const std::shared_ptr &module, + const PipelineStageLayout &stage) { + if (!module) { + throw std::invalid_argument("pipeline module must not be null"); + } + + if (stage.has_final_norm != stage.has_lm_head) { + throw std::invalid_argument("final norm and LM head must belong to the same pipeline stage"); + } + + std::vector> chunks; + chunks.reserve(stage.chunks.size()); + for (std::size_t chunk_idx = 0; chunk_idx < stage.chunks.size(); ++chunk_idx) { + const PipelineChunkLayout &chunk_layout = stage.chunks[chunk_idx]; + if (chunk_layout.stage_id != stage.stage_id || chunk_layout.global_chunk_id < 0) { + throw std::invalid_argument("pipeline chunk identity is inconsistent with its stage"); + } + if (chunk_layout.local_chunk_idx != static_cast(chunk_idx)) { + throw std::invalid_argument("pipeline local chunk indices must be contiguous and start at zero"); + } + + std::vector> chunk_parts; + if (chunk_idx == 0 && stage.has_embedding) { + chunk_parts.push_back(module->mutable_module(Module::kPPFirstStageName)); + } + chunk_parts.push_back(module->mutable_module(std::string(Module::kPPChunkNamePrefix) + + std::to_string(chunk_layout.local_chunk_idx))); + if (chunk_idx == stage.chunks.size() - 1 && stage.has_final_norm) { + chunk_parts.push_back(module->mutable_module(Module::kPPLastStageName)); + } + chunks.push_back(std::make_shared(std::move(chunk_parts))); + } + return chunks; +} } // namespace thread_local int pp_rank = 0; @@ -23,7 +63,8 @@ void PipelineParallel::BuildPipelineStage(const std::vector } void PipelineParallel::SetupSchedule(int num_micro_batches) { - schedule_ = std::make_shared(pipeline_stage_, num_stages_, num_micro_batches); + schedule_ = std::make_shared(pipeline_stage_, num_stages_, num_micro_batches, pipeline_layout_, + explicit_chunk_mode_); } float PipelineParallel::TrainStep(const std::vector> &input, @@ -31,49 +72,21 @@ float PipelineParallel::TrainStep(const std::vector> &in const std::shared_ptr &optimizer, const std::shared_ptr &loss_fn, DataType dtype) { std::shared_ptr stage_input; - std::shared_ptr stage_target = target[0]; - if (rank_ == 0) { - stage_input = input[0]; + if (rank_ == (pipeline_layout_ ? pipeline_layout_->GetInputStage() : 0)) { + stage_input = input.at(0); } - return schedule_->Step(stage_input, stage_target, optimizer, loss_fn, dtype); + return schedule_->Step(stage_input, target.at(0), optimizer, loss_fn, dtype); } StageInfo PipelineParallel::GetStageInfo(int total_layers, int pp_size, int rank, int chunks_per_stage) { - bool is_first_stage = (rank == 0); - bool is_last_stage = (rank == pp_size - 1); - - std::vector> layer_ranges_per_chunk; - - int layers_per_chunk = total_layers / (pp_size * chunks_per_stage); - int remainder = total_layers % (pp_size * chunks_per_stage); - - for (int local_chunk_idx = 0; local_chunk_idx < chunks_per_stage; ++local_chunk_idx) { - int global_chunk_idx = local_chunk_idx * pp_size + rank; - - if (global_chunk_idx * layers_per_chunk >= total_layers) { - break; - } - - int chunk_start = global_chunk_idx * layers_per_chunk; - int chunk_end = chunk_start + layers_per_chunk; - - if (global_chunk_idx < remainder) { - // Assign an additional layer to each of the first remainder chunks - chunk_start = global_chunk_idx * (layers_per_chunk + 1); - chunk_end = chunk_start + (layers_per_chunk + 1); - } else { - chunk_start = remainder * (layers_per_chunk + 1) + (global_chunk_idx - remainder) * layers_per_chunk; - chunk_end = chunk_start + layers_per_chunk; - } - - chunk_end = std::min(chunk_end, total_layers); - if (chunk_start < chunk_end) { - layer_ranges_per_chunk.push_back({chunk_start, chunk_end}); - } + const auto layout = PipelineLayout::BuildUniformLayout(total_layers, pp_size, chunks_per_stage); + const auto &stage = layout.GetStage(rank); + std::vector> local_ranges; + for (const auto &chunk : stage.chunks) { + local_ranges.emplace_back(static_cast(chunk.layer_range.begin), static_cast(chunk.layer_range.end)); } - - return {is_first_stage, is_last_stage, layer_ranges_per_chunk}; + return {rank == 0, rank == pp_size - 1, std::move(local_ranges)}; } PipelineParallel::PipelineParallel(const std::shared_ptr module, int num_stages, int num_micro_batches, @@ -102,6 +115,29 @@ PipelineParallel::PipelineParallel(const std::shared_ptr module, int num SetupSchedule(num_micro_batches); } +PipelineParallel::PipelineParallel(const std::shared_ptr module, int num_stages, int num_micro_batches, + const std::vector> &recv_shape, int pp_rank, Device device, + std::shared_ptr pipeline_layout, bool explicit_chunk_mode) + : num_stages_(num_stages), rank_(pp_rank), pipeline_layout_(std::move(pipeline_layout)), + explicit_chunk_mode_(explicit_chunk_mode) { + if (pipeline_layout_ == nullptr) { + throw std::invalid_argument("pipeline layout must not be null"); + } + if (pipeline_layout_->GetNumStages() != num_stages_) { + throw std::invalid_argument("pipeline layout stage count does not match PipelineParallel::num_stages"); + } + + const PipelineStageLayout &stage = pipeline_layout_->GetStage(rank_); + if (const auto transformer = std::dynamic_pointer_cast(module)) { + if (transformer->GetPipelineLayout() != pipeline_layout_ || transformer->GetStageId() != rank_) { + throw std::invalid_argument("pipeline wrapper must share the model's layout and stage"); + } + } + modules_[kModuleName] = module; + auto chunks = BuildPipelineChunksFromLayout(module, stage); + BuildPipelineStage(recv_shape, device, std::move(chunks)); + SetupSchedule(num_micro_batches); +} std::vector> *PipelineParallel::mutable_chunks() { return pipeline_stage_->mutable_chunks(); } } // namespace infini_train::nn::parallel diff --git a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc index 6578e628b..5b536905e 100644 --- a/infini_train/src/nn/parallel/pp/pipeline_schedule.cc +++ b/infini_train/src/nn/parallel/pp/pipeline_schedule.cc @@ -1,8 +1,12 @@ // pipeline_schedule.cc #include "infini_train/include/nn/parallel/pp/pipeline_schedule.h" +#include #include +#include #include +#include +#include #include #include "glog/logging.h" @@ -13,6 +17,7 @@ #include "infini_train/include/nn/init.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" #include "infini_train/include/nn/parallel/pp/pipeline_stage.h" #include "infini_train/include/nn/parallel/pp/send_recv.h" #include "infini_train/include/optimizer.h" @@ -20,26 +25,41 @@ namespace infini_train::nn::parallel { -void PrintScheduleTable(const std::vector &schedule, int n, int num_stages, - int vpp_size) { - int total_global_chunks = num_stages * vpp_size; - - LOG(INFO) << std::format("=== Schedule Table ===\n" - "n: {}, stages: {}, vpp: {}, total_chunks: {}", - n, num_stages, vpp_size, total_global_chunks); - LOG(INFO) << ""; - LOG(INFO) << "Step | Type | Microbatch | Global Chunk | Local Chunk | Stage"; - LOG(INFO) << "-----|-----------|------------|--------------|-------------|-------"; +PipelineSchedule::PipelineSchedule(std::shared_ptr stage, int num_stages, int num_micro_batches, + std::shared_ptr layout, bool explicit_chunk_mode) + : num_micro_batches_(num_micro_batches), stage_(std::move(stage)), layout_(std::move(layout)), + explicit_chunk_mode_(explicit_chunk_mode) { + if (!stage_ || num_micro_batches_ <= 0 || num_stages != stage_->num_stages()) { + throw std::invalid_argument("invalid pipeline stage or microbatch count"); + } + if (explicit_chunk_mode_ && !layout_) { + throw std::invalid_argument("explicit execution requires a layout"); + } + if (layout_ + && (layout_->GetNumStages() != num_stages + || stage_->chunks().size() != layout_->GetStage(stage_->stage_index()).chunks.size())) { + throw std::invalid_argument("stage modules disagree with pipeline layout"); + } + if (layout_ && !explicit_chunk_mode_) { + const int count = layout_->GetNumChunks(); + if (count <= 0 || count % num_stages != 0) { + throw std::invalid_argument("ordinary execution requires a full interleaved layout"); + } + for (int gid = 0; gid < count; ++gid) { + const auto &chunk = layout_->GetChunk(gid); + if (chunk.stage_id != gid % num_stages || chunk.local_chunk_idx != gid / num_stages) { + throw std::invalid_argument("explicit ownership requires explicit execution mode"); + } + } + } +} +void PrintScheduleTable(const std::vector &schedule, int n, int total_chunks) { + LOG(INFO) << "Pipeline schedule: microbatches=" << n << ", chunks=" << total_chunks; + LOG(INFO) << "Step | Type | Microbatch | Global Chunk | Local Chunk | Stage"; for (const auto &task : schedule) { - int owning_stage = task.global_chunk_id % num_stages; - int local_chunk = task.global_chunk_id / num_stages; - - std::string type_str = task.is_forward ? "Forward" : "Backward"; - - auto s_info = std::format("{:4} | {:<9} | {:>10} | {:>12} | {:>11} | {:>5}", task.step, type_str, - task.microbatch_id, task.global_chunk_id, local_chunk, owning_stage); - LOG(INFO) << s_info; + LOG(INFO) << task.step << " | " << (task.is_forward ? "Forward" : "Backward") << " | " << task.microbatch_id + << " | " << task.global_chunk_id << " | " << task.local_chunk_idx << " | " << task.stage_id; } } @@ -69,24 +89,98 @@ std::vector> PipelineSchedule::SendToNext(const std::vec } PipelineParallelScheduler::Task PipelineParallelScheduler::CreateTask(int step, int mb, int global_chunk, - int num_stages, int total_chunks, - bool is_forward) { + int num_stages, int total_chunks, bool is_forward, + const PipelineLayout *layout) { PipelineParallelScheduler::Task task; task.step = step; task.microbatch_id = mb; task.global_chunk_id = global_chunk; - task.local_chunk_idx = global_chunk / num_stages; + task.local_chunk_idx = layout ? layout->GetChunk(global_chunk).local_chunk_idx : global_chunk / num_stages; task.is_forward = is_forward; - task.stage_id = global_chunk % num_stages; + task.stage_id = layout ? layout->GetChunk(global_chunk).stage_id : global_chunk % num_stages; task.is_last_chunk = (global_chunk == total_chunks - 1); task.is_first_chunk = (global_chunk == 0); return task; } -std::vector PipelineParallelScheduler::GenerateGPipeSchedule(int n, int num_stages, - int vpp_size) { +std::vector +PipelineParallelScheduler::GenerateGPipeSchedule(int n, int num_stages, int vpp_size, const PipelineLayout *layout) { std::vector schedule; - int total_global_chunks = num_stages * vpp_size; + if (n <= 0) { + return schedule; + } + if (num_stages <= 0 || vpp_size <= 0 + || static_cast(num_stages) * vpp_size > std::numeric_limits::max()) { + throw std::invalid_argument("invalid schedule dimensions"); + } + int total_global_chunks = layout ? layout->GetNumChunks() : num_stages * vpp_size; + if (total_global_chunks <= 0 || total_global_chunks > std::numeric_limits::max() / 2 + || n > std::numeric_limits::max() / 2 - total_global_chunks) { + throw std::invalid_argument("schedule step count overflow"); + } + + bool interleaved = true; + if (layout) { + if (layout->GetNumStages() != num_stages || layout->GetMaxLocalChunks() != vpp_size) { + throw std::invalid_argument("schedule dimensions disagree with layout"); + } + interleaved = (total_global_chunks == static_cast(num_stages) * vpp_size); + for (int gid = 0; interleaved && gid < total_global_chunks; ++gid) { + const auto &chunk = layout->GetChunk(gid); + interleaved = (chunk.stage_id == gid % num_stages) && (chunk.local_chunk_idx == gid / num_stages); + } + } + // A rank revisited after a cross-rank edge cannot safely pipeline the old + // diagonal task order: Send/Recv put a completion wait on its compute stream. + // Example: owners 0,1,0 can make rank 0 send mb1 while rank 1 sends mb0 back. + std::vector visited(num_stages, false); + int previous_owner = -1; + bool revisits_stage = false; + for (int gid = 0; gid < total_global_chunks; ++gid) { + const int owner = layout ? layout->GetChunk(gid).stage_id : gid % num_stages; + if (owner != previous_owner) { + revisits_stage = revisits_stage || visited[owner]; + visited[owner] = true; + previous_owner = owner; + } + } + if (revisits_stage) { + // Here step numbers tasks, not diagonal waves; protect 2*n*chunks. + if (n > std::numeric_limits::max() / 2 / total_global_chunks) { + throw std::invalid_argument("schedule task count overflow"); + } + int step = 0; + for (int mb = 0; mb < n; ++mb) { + for (int gid = 0; gid < total_global_chunks; ++gid) { + schedule.push_back(CreateTask(step++, mb, gid, num_stages, total_global_chunks, true, layout)); + } + } + for (int mb = n - 1; mb >= 0; --mb) { + for (int gid = total_global_chunks - 1; gid >= 0; --gid) { + schedule.push_back(CreateTask(step++, mb, gid, num_stages, total_global_chunks, false, layout)); + } + } + return schedule; + } + if (!interleaved) { + const int waves = n + total_global_chunks - 1; + for (int direction = 0; direction < 2; ++direction) { + for (int wave = 0; wave < waves; ++wave) { + for (int mb = 0; mb < n; ++mb) { + const int offset = wave - mb; + if (offset < 0 || offset >= total_global_chunks) { + continue; + } + const int gid = direction == 0 ? offset : total_global_chunks - 1 - offset; + schedule.push_back(CreateTask(wave + direction * waves, mb, gid, num_stages, total_global_chunks, + direction == 0, layout)); + } + } + } + return schedule; + } + + // Keep the existing GPipe wave order for layouts without rank revisits int total_steps = n + total_global_chunks - 1; // ======== Forward Pass ======== @@ -95,7 +189,7 @@ std::vector PipelineParallelScheduler::Generate int global_chunk_id = step - mb; if (global_chunk_id >= 0 && global_chunk_id < total_global_chunks) { auto is_forward = true; - auto task = CreateTask(step, mb, global_chunk_id, num_stages, total_global_chunks, is_forward); + auto task = CreateTask(step, mb, global_chunk_id, num_stages, total_global_chunks, is_forward, layout); schedule.push_back(task); } } @@ -107,8 +201,8 @@ std::vector PipelineParallelScheduler::Generate int global_chunk_id = (total_steps - 1 - step) - mb; if (global_chunk_id >= 0 && global_chunk_id < total_global_chunks) { auto is_forward = false; - auto task - = CreateTask(step + total_steps, mb, global_chunk_id, num_stages, total_global_chunks, is_forward); + auto task = CreateTask(step + total_steps, mb, global_chunk_id, num_stages, total_global_chunks, + is_forward, layout); schedule.push_back(task); } } @@ -127,14 +221,23 @@ std::vector PipelineParallelScheduler::Generate } std::vector -PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages, int vpp_size) { +PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages, int vpp_size, + const PipelineLayout *layout) { std::vector schedule; - + // Preserve legacy empty-schedule behavior for non-positive dimensions. if (n <= 0 || num_stages <= 0 || vpp_size <= 0) { return schedule; } + if (n <= 0 || num_stages <= 0 || vpp_size <= 0 + || static_cast(num_stages) * vpp_size > std::numeric_limits::max()) { + throw std::invalid_argument("invalid schedule dimension"); + } - int total_global_chunks = num_stages * vpp_size; + int total_global_chunks = layout ? layout->GetNumChunks() : num_stages * vpp_size; + if (total_global_chunks <= 0 || total_global_chunks > std::numeric_limits::max() / 2 + || n > std::numeric_limits::max() / 2 - total_global_chunks) { + throw std::invalid_argument("schedule step count overflow"); + } int warmup_steps = total_global_chunks - 1; int total_steps = 2 * warmup_steps + n; @@ -145,7 +248,8 @@ PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages int forward_global_chunk = step - mb; if (forward_global_chunk >= 0 && forward_global_chunk < total_global_chunks) { auto is_forward = true; - auto task = CreateTask(step, mb, forward_global_chunk, num_stages, total_global_chunks, is_forward); + auto task + = CreateTask(step, mb, forward_global_chunk, num_stages, total_global_chunks, is_forward, layout); schedule.push_back(task); } } @@ -160,7 +264,8 @@ PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages int forward_global_chunk = step - mb; if (forward_global_chunk >= 0 && forward_global_chunk < total_global_chunks) { auto is_forward = true; - auto task = CreateTask(step, mb, forward_global_chunk, num_stages, total_global_chunks, is_forward); + auto task + = CreateTask(step, mb, forward_global_chunk, num_stages, total_global_chunks, is_forward, layout); schedule.push_back(task); } @@ -169,7 +274,8 @@ PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages if (backward_global_chunk >= 0 && backward_global_chunk < total_global_chunks) { auto is_forward = false; - auto task = CreateTask(step, mb, backward_global_chunk, num_stages, total_global_chunks, is_forward); + auto task + = CreateTask(step, mb, backward_global_chunk, num_stages, total_global_chunks, is_forward, layout); schedule.push_back(task); } } @@ -182,7 +288,8 @@ PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages int backward_global_chunk = (total_global_chunks - 1) - (backward_step - mb); if (backward_global_chunk >= 0 && backward_global_chunk < total_global_chunks) { auto is_forward = false; - auto task = CreateTask(step, mb, backward_global_chunk, num_stages, total_global_chunks, is_forward); + auto task + = CreateTask(step, mb, backward_global_chunk, num_stages, total_global_chunks, is_forward, layout); schedule.push_back(task); } } @@ -194,105 +301,131 @@ PipelineParallelScheduler::GenerateInterleaved1F1BSchedule(int n, int num_stages float PipelineSchedule::StepMicroBatches(const std::vector> µbatch_inputs, const std::vector> µbatch_targets, const std::shared_ptr &loss_fn, DataType dtype) { - int n = num_micro_batches_; + + const int n = num_micro_batches_; int num_stages = stage_->num_stages(); int stage_idx = stage_->stage_index(); - int vpp_size = global::GetVirtualPipelineParallelSize(); - - auto schedule = PipelineParallelScheduler::GenerateGPipeSchedule(n, num_stages, vpp_size); - - static bool has_printed = false; - if (!has_printed && stage_idx == 0) { - PrintScheduleTable(schedule, n, num_stages, vpp_size); - has_printed = true; + int vpp_size = layout_ ? layout_->GetMaxLocalChunks() : global::GetVirtualPipelineParallelSize(); + const auto local_count = stage_->chunks().size(); + if (microbatch_inputs.size() != static_cast(n) + || microbatch_targets.size() != static_cast(n)) { + throw std::invalid_argument("microbatch containers disagree with microbatch count"); } - float total_loss = 0.0f; - - std::vector>>> activations( - vpp_size, std::vector>>(n)); + auto schedule = PipelineParallelScheduler::GenerateGPipeSchedule(n, num_stages, vpp_size, layout_.get()); + if (!has_printed_schedule_ && stage_idx == (layout_ ? layout_->GetInputStage() : 0)) { + PrintScheduleTable(schedule, n, layout_ ? layout_->GetNumChunks() : num_stages * vpp_size); + has_printed_schedule_ = true; + } + using TensorList = std::vector>; + std::vector> activations(local_count, std::vector(n)); std::vector> no_sync_guards; - no_sync_guards.reserve(stage_->chunks().size()); for (const auto &chunk : stage_->chunks()) { no_sync_guards.push_back(chunk->no_sync()); } - std::vector backward_counts(vpp_size, 0); + std::vector backward_counts(local_count, 0); + float total_loss = 0.0f; - for (size_t i = 0; i < schedule.size(); ++i) { - const auto &task = schedule[i]; + for (const auto &task : schedule) { if (task.stage_id != stage_idx) { continue; } - - int mb = task.microbatch_id; + const int gid = task.global_chunk_id; + const int local = task.local_chunk_idx; + const int mb = task.microbatch_id; if (task.is_forward) { infini_train::AutocastGuard autocast_guard(stage_->device().type(), dtype); - - std::vector> inputs; - + TensorList inputs; if (task.is_first_chunk) { - inputs = {microbatch_inputs[mb]}; - } else { - if (stage_->IsFirstStage()) { - inputs = ReceiveFromPrev(num_stages - 1); - } else { - inputs = ReceiveFromPrev(stage_->prev_rank()); + if (!microbatch_inputs.at(mb)) { + throw std::invalid_argument("missing input on logical input stage"); } + inputs = {microbatch_inputs.at(mb)}; + } else if (layout_) { + const auto &previous = layout_->GetChunk(gid - 1); + inputs = previous.stage_id == stage_idx ? activations.at(previous.local_chunk_idx).at(mb) + : ReceiveFromPrev(previous.stage_id); + } else { + inputs = ReceiveFromPrev(stage_->IsFirstStage() ? num_stages - 1 : stage_->prev_rank()); } - - activations[task.local_chunk_idx][mb] = stage_->ForwardOneChunk(inputs, task.local_chunk_idx); - + auto &output = activations.at(local).at(mb); + output = stage_->ForwardOneChunk(inputs, local); if (!task.is_last_chunk) { - if (stage_->IsLastStage()) { - SendToNext(activations[task.local_chunk_idx][mb], 0); - } else { - SendToNext(activations[task.local_chunk_idx][mb], stage_->next_rank()); + const int next_owner + = layout_ ? layout_->GetChunk(gid + 1).stage_id : (stage_->IsLastStage() ? 0 : stage_->next_rank()); + if (!layout_ || next_owner != stage_idx) { + output = SendToNext(output, next_owner); } } } else { - const bool is_last_microbatch = ++backward_counts[task.local_chunk_idx] == n; - if (is_last_microbatch) { - no_sync_guards[task.local_chunk_idx].reset(); + // A successor on the same stage already traversed this local graph. + if (layout_ && !task.is_last_chunk && layout_->GetChunk(gid + 1).stage_id == stage_idx) { + continue; + } + if (layout_) { + for (int chain_gid = gid; chain_gid >= 0; --chain_gid) { + const auto &chunk = layout_->GetChunk(chain_gid); + if (chunk.stage_id != stage_idx) { + break; + } + if (++backward_counts.at(chunk.local_chunk_idx) == n) { + no_sync_guards.at(chunk.local_chunk_idx).reset(); + } + } + } else if (++backward_counts.at(local) == n) { + no_sync_guards.at(local).reset(); + } + const auto &output = activations.at(local).at(mb); + if (output.empty() || !output[0]) { + throw std::invalid_argument("missing activation "); } if (task.is_last_chunk) { - auto target = microbatch_targets[mb]; + if (!microbatch_targets.at(mb) || !loss_fn) { + throw std::invalid_argument("missing target or loss function"); + } std::shared_ptr loss; { infini_train::AutocastGuard autocast_guard(stage_->device().type(), dtype); - - auto target_on_device = target->To(activations[task.local_chunk_idx][mb][0]->GetDevice()); - loss = (*loss_fn)( - {activations[task.local_chunk_idx][mb][0], std::make_shared(target_on_device)})[0]; - loss = loss / n; + auto target = microbatch_targets.at(mb)->To(output[0]->GetDevice()); + loss = (*loss_fn)({output[0], std::make_shared(target)})[0] / n; } loss->Backward(); - // Defer the loss D2H copy until after backward; reading it earlier would synchronize CUDA - // between forward and backward. total_loss += static_cast(loss->To(Device()).DataPtr())[0]; } else { - auto out_tensor = activations[task.local_chunk_idx][mb][0]; - - auto dummy_gradient - = std::make_shared(out_tensor->Dims(), out_tensor->Dtype(), out_tensor->GetDevice()); - - out_tensor->Backward(dummy_gradient); + auto dummy = std::make_shared(output[0]->Dims(), output[0]->Dtype(), output[0]->GetDevice()); + dummy->Fill(0.0f); + output[0]->Backward(dummy); } } } - + for (int completed : backward_counts) { + if (completed != n) { + throw std::runtime_error("incomplete local chunk backward traversal"); + } + } return total_loss; } float PipelineSchedule::Step(std::shared_ptr input, std::shared_ptr target, const std::shared_ptr &optimizer, const std::shared_ptr &loss_fn, DataType dtype) { - std::vector> micro_batches(num_micro_batches_); - std::vector> target_mbs(num_micro_batches_); - if (stage_->IsFirstStage()) { - micro_batches = input->Split(input->Dims()[0] / num_micro_batches_); - } + const int stage_idx = stage_->stage_index(); + const int n = num_micro_batches_; + const int input_owner = layout_ ? layout_->GetInputStage() : 0; + const int output_owner = layout_ ? layout_->GetOutputStage() : stage_->num_stages() - 1; + auto split = [this, n](const std::shared_ptr &tensor) { + if (explicit_chunk_mode_ + && (!tensor || tensor->Dims().empty() || tensor->Dims()[0] <= 0 || tensor->Dims()[0] % n != 0)) { + throw std::invalid_argument("invalid input/target microbatch dimension"); + } + return tensor->Split(tensor->Dims()[0] / n); + }; + std::vector> micro_batches(n), target_mbs(n); - if (stage_->IsLastStage()) { - target_mbs = target->Split(target->Dims()[0] / num_micro_batches_); + if (stage_idx == input_owner) { + micro_batches = split(input); + } + if (stage_idx == output_owner) { + target_mbs = split(target); } optimizer->ZeroGrad(); diff --git a/infini_train/src/utils/precision_checker.cc b/infini_train/src/utils/precision_checker.cc index 2965284eb..1e5d4f436 100644 --- a/infini_train/src/utils/precision_checker.cc +++ b/infini_train/src/utils/precision_checker.cc @@ -1,5 +1,6 @@ #include "infini_train/include/utils/precision_checker.h" +#include #include #include #include @@ -11,10 +12,14 @@ #include #include #include +#include +#include +#include #include "infini_train/include/autograd/function.h" #include "infini_train/include/nn/modules/module.h" #include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" #include "infini_train/include/nn/parallel/tensor_parallel.h" #include "infini_train/include/tensor.h" #include "infini_train/include/utils/global_module_hook_registry.h" @@ -337,10 +342,9 @@ void PrecisionChecker::CheckTensors(const std::string &stage, const std::string // Original precision MD5 md5 = ComputeMD5(cpu_tensor->DataPtr(), byte_size); } - log_stream << context_key << " " << log_name << " tensor[" << i << "]: " - << "dtype=" << DataTypeToString(cpu_tensor->Dtype()) << " " - << "shape=" << FormatShape(cpu_tensor->Dims()) << " " - << "md5=" << md5 << std::endl; + log_stream << context_key << " " << log_name << " tensor[" << i + << "]: " << "dtype=" << DataTypeToString(cpu_tensor->Dtype()) << " " + << "shape=" << FormatShape(cpu_tensor->Dims()) << " " << "md5=" << md5 << std::endl; } else { // Simple format (default) TensorStats stats = ComputeStats(float_data, num_elements); @@ -349,12 +353,10 @@ void PrecisionChecker::CheckTensors(const std::string &stage, const std::string = (config.check_nan && stats.nan_count > 0) || (config.check_inf && stats.inf_count > 0); const std::string error_marker = has_error ? " <- ERROR" : ""; - log_stream << context_key << " " << log_name << " tensor[" << i << "]: " - << "dtype=" << DataTypeToString(cpu_tensor->Dtype()) << " " - << "shape=" << FormatShape(cpu_tensor->Dims()) << " " - << "min=" << stats.min_val << " " - << "max=" << stats.max_val << " " - << "mean=" << stats.mean_val << " ["; + log_stream << context_key << " " << log_name << " tensor[" << i + << "]: " << "dtype=" << DataTypeToString(cpu_tensor->Dtype()) << " " + << "shape=" << FormatShape(cpu_tensor->Dims()) << " " << "min=" << stats.min_val << " " + << "max=" << stats.max_val << " " << "mean=" << stats.mean_val << " ["; // Print first 6 values constexpr size_t max_print = 6; @@ -406,7 +408,34 @@ static inline bool ShouldSkipNameMap(std::string_view name) { return name.rfind("__pp", 0) == 0; // starts_with("__pp") } -void PrecisionChecker::BuildNameMap(nn::Module *root_model) { +static std::string ToGlobalModuleName(std::string_view name, + const std::shared_ptr &layout, int pp_rank) { + + constexpr std::string_view prefix = "transformer.h."; + if (!layout || !name.starts_with(prefix)) { + return std::string(name); + } + const auto &stage = layout->GetStage(pp_rank); + const auto tail = name.substr(prefix.size()); + const auto dot = tail.find('.'); + const auto token = tail.substr(0, dot); + nn::parallel::LayerIndex local = -1; + const auto [end, error] = std::from_chars(token.data(), token.data() + token.size(), local); + if (token.empty() || error != std::errc{} || end != token.data() + token.size() || local < 0) { + throw std::out_of_range("invalid local layer in precision name: " + std::string(name)); + } + for (const auto &chunk : stage.chunks) { + if (local < chunk.layer_range.Size()) { + return std::string(prefix) + std::to_string(chunk.layer_range.begin + local) + + (dot == std::string_view::npos ? "" : std::string(tail.substr(dot))); + } + local -= chunk.layer_range.Size(); + } + throw std::out_of_range("local layer is not owned by stage: " + std::string(name)); +} + +void PrecisionChecker::BuildNameMap(nn::Module *root_model, std::shared_ptr layout, + int pp_rank) { const auto &global_config = PrecisionCheckEnv::Instance().GetConfig(); if (global_config.level == PrecisionCheckLevel::OFF || root_model == nullptr) { return; @@ -423,7 +452,7 @@ void PrecisionChecker::BuildNameMap(nn::Module *root_model) { if (ShouldSkipNameMap(name)) { continue; // skip PP internal tree } - g_module_name_map[module.get()] = name; // keep InfiniTrain path directly + g_module_name_map[module.get()] = ToGlobalModuleName(name, layout, pp_rank); } } diff --git a/tests/distributed/CMakeLists.txt b/tests/distributed/CMakeLists.txt index b8ed49700..bf8cf2628 100644 --- a/tests/distributed/CMakeLists.txt +++ b/tests/distributed/CMakeLists.txt @@ -23,3 +23,44 @@ set_tests_properties(RankTest.MultiNodeSingleProcessIsParallel LABELS cpu TIMEOUT 10 ) + +infini_train_add_test(test_pipeline_layout + SOURCES test_pipeline_layout.cc + LABELS cpu +) + +infini_train_add_test(test_pipeline_parallel + SOURCES test_pipeline_parallel.cc + LABELS cpu +) + +infini_train_add_test(test_pp_layout_parser + SOURCES test_pp_layout_parser.cc ${CMAKE_SOURCE_DIR}/example/common/parser.cc + LABELS cpu +) + +# Loader fixtures run in separate CPU processes because GlobalEnv is initialized once. +add_executable(test_pipeline_loader + test_pipeline_loader.cc + ${CMAKE_SOURCE_DIR}/example/common/parser.cc + ${CMAKE_SOURCE_DIR}/example/common/utils.cc + ${CMAKE_SOURCE_DIR}/example/gpt2/checkpoint_loader.cc + ${CMAKE_SOURCE_DIR}/example/llama3/checkpoint_loader.cc +) +target_link_libraries(test_pipeline_loader PRIVATE GTest::gtest) +link_infini_train_exe(test_pipeline_loader) +foreach(topology IN ITEMS "1:1" "2:1" "2:2") + string(REPLACE ":" ";" dimensions "${topology}") + list(GET dimensions 0 pp) + list(GET dimensions 1 vpp) + add_test(NAME PipelineLoader.PP${pp}VPP${vpp} + COMMAND ${CMAKE_COMMAND} -E env WORLD_SIZE=${pp} LOCAL_WORLD_SIZE=${pp} RANK=0 LOCAL_RANK=0 + PIPELINE_LOADER_TEST_PP=${pp} PIPELINE_LOADER_TEST_VPP=${vpp} + $) + set_tests_properties(PipelineLoader.PP${pp}VPP${vpp} PROPERTIES LABELS cpu TIMEOUT 30) +endforeach() + +infini_train_add_test(test_pipeline_schedule + SOURCES test_pipeline_schedule.cc + LABELS cpu +) diff --git a/tests/distributed/test_pipeline_layout.cc b/tests/distributed/test_pipeline_layout.cc new file mode 100644 index 000000000..9d2c2b8b0 --- /dev/null +++ b/tests/distributed/test_pipeline_layout.cc @@ -0,0 +1,315 @@ +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include + +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" + +namespace infini_train::nn::parallel { +namespace { + +static_assert(!std::is_default_constructible_v); + +void ExpectRange(const LayerRange &range, LayerIndex expected_begin, LayerIndex expected_end) { + EXPECT_EQ(range.begin, expected_begin); + EXPECT_EQ(range.end, expected_end); +} + +TEST(LayerRangeTest, UsesHalfOpenBounds) { + const LayerRange range{4, 12}; + + EXPECT_EQ(range.Size(), 8); + EXPECT_FALSE(range.Contains(3)); + EXPECT_TRUE(range.Contains(4)); + EXPECT_TRUE(range.Contains(11)); + EXPECT_FALSE(range.Contains(12)); +} + +TEST(PipelineLayoutTest, CreatesUniformLayouts) { + const PipelineLayout even = PipelineLayout::BuildUniformLayout(12, 2); + EXPECT_EQ(even.GetNumLayers(), 12); + EXPECT_EQ(even.GetNumStages(), 2); + EXPECT_FALSE(even.IsCustom()); + + const PipelineStageLayout &stage0 = even.GetStage(0); + const PipelineStageLayout &stage1 = even.GetStage(1); + ASSERT_EQ(stage0.chunks.size(), 1U); + ASSERT_EQ(stage1.chunks.size(), 1U); + EXPECT_EQ(stage0.stage_id, 0); + EXPECT_EQ(stage1.stage_id, 1); + EXPECT_EQ(stage0.chunks[0].local_chunk_idx, 0); + EXPECT_EQ(stage1.chunks[0].local_chunk_idx, 0); + ExpectRange(stage0.chunks[0].layer_range, 0, 6); + ExpectRange(stage1.chunks[0].layer_range, 6, 12); + EXPECT_TRUE(stage0.has_embedding); + EXPECT_FALSE(stage0.has_final_norm); + EXPECT_FALSE(stage0.has_lm_head); + EXPECT_FALSE(stage1.has_embedding); + EXPECT_TRUE(stage1.has_final_norm); + EXPECT_TRUE(stage1.has_lm_head); + + const PipelineLayout uneven = PipelineLayout::BuildUniformLayout(13, 2); + ASSERT_EQ(uneven.GetStage(0).chunks.size(), 1U); + ASSERT_EQ(uneven.GetStage(1).chunks.size(), 1U); + ExpectRange(uneven.GetStage(0).chunks[0].layer_range, 0, 7); + ExpectRange(uneven.GetStage(1).chunks[0].layer_range, 7, 13); + + const PipelineLayout single = PipelineLayout::BuildUniformLayout(12, 1); + const PipelineStageLayout &single_stage = single.GetStage(0); + ASSERT_EQ(single_stage.chunks.size(), 1U); + ExpectRange(single_stage.chunks[0].layer_range, 0, 12); + EXPECT_TRUE(single_stage.has_embedding); + EXPECT_TRUE(single_stage.has_final_norm); + EXPECT_TRUE(single_stage.has_lm_head); +} + +TEST(PipelineLayoutTest, OneLayerPerStage) { + const auto layout = PipelineLayout::BuildUniformLayout(4, 4); + for (int stage = 0; stage < 4; ++stage) { + const auto &range = layout.GetStage(stage).chunks.at(0).layer_range; + EXPECT_EQ(range.begin, stage); + EXPECT_EQ(range.end, stage + 1); + EXPECT_EQ(layout.GetLocalLayerIndex(stage, stage), 0); + } +} + +TEST(PipelineLayoutInvalidTest, RejectsInvalidUniformArguments) { + EXPECT_THROW((void)PipelineLayout::BuildUniformLayout(-1, 2), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildUniformLayout(12, 0), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildUniformLayout(12, -1), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildUniformLayout(12, 2, 0), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildUniformLayout(12, 2, std::numeric_limits::max()), + std::invalid_argument); +} + +TEST(PipelineLayoutTest, CreatesCustomContinuousLayouts) { + const std::vector counts{4, 8, 6, 6}; + const PipelineLayout layout = PipelineLayout::BuildCustomLayout(24, 4, counts); + + EXPECT_EQ(layout.GetNumLayers(), 24); + EXPECT_EQ(layout.GetNumStages(), 4); + EXPECT_TRUE(layout.IsCustom()); + + const PipelineStageLayout &stage0 = layout.GetStage(0); + const PipelineStageLayout &stage1 = layout.GetStage(1); + const PipelineStageLayout &stage2 = layout.GetStage(2); + const PipelineStageLayout &stage3 = layout.GetStage(3); + ASSERT_EQ(stage0.chunks.size(), 1U); + ASSERT_EQ(stage1.chunks.size(), 1U); + ASSERT_EQ(stage2.chunks.size(), 1U); + ASSERT_EQ(stage3.chunks.size(), 1U); + EXPECT_EQ(stage0.stage_id, 0); + EXPECT_EQ(stage1.stage_id, 1); + EXPECT_EQ(stage2.stage_id, 2); + EXPECT_EQ(stage3.stage_id, 3); + EXPECT_EQ(stage0.chunks[0].local_chunk_idx, 0); + EXPECT_EQ(stage1.chunks[0].local_chunk_idx, 0); + EXPECT_EQ(stage2.chunks[0].local_chunk_idx, 0); + EXPECT_EQ(stage3.chunks[0].local_chunk_idx, 0); + ExpectRange(stage0.chunks[0].layer_range, 0, 4); + ExpectRange(stage1.chunks[0].layer_range, 4, 12); + ExpectRange(stage2.chunks[0].layer_range, 12, 18); + ExpectRange(stage3.chunks[0].layer_range, 18, 24); + + EXPECT_TRUE(stage0.has_embedding); + EXPECT_FALSE(stage1.has_embedding); + EXPECT_FALSE(stage2.has_embedding); + EXPECT_FALSE(stage3.has_embedding); + EXPECT_FALSE(stage0.has_final_norm); + EXPECT_FALSE(stage1.has_final_norm); + EXPECT_FALSE(stage2.has_final_norm); + EXPECT_TRUE(stage3.has_final_norm); + EXPECT_FALSE(stage0.has_lm_head); + EXPECT_FALSE(stage1.has_lm_head); + EXPECT_FALSE(stage2.has_lm_head); + EXPECT_TRUE(stage3.has_lm_head); + + const PipelineLayout gpt2_layout = [] { + const std::vector gpt2_counts{7, 5}; + return PipelineLayout::BuildCustomLayout(12, 2, gpt2_counts); + }(); + ASSERT_EQ(gpt2_layout.GetStage(0).chunks.size(), 1U); + ASSERT_EQ(gpt2_layout.GetStage(1).chunks.size(), 1U); + ExpectRange(gpt2_layout.GetStage(0).chunks[0].layer_range, 0, 7); + ExpectRange(gpt2_layout.GetStage(1).chunks[0].layer_range, 7, 12); +} + +TEST(PipelineLayoutInvalidTest, RejectsInvalidCustomArguments) { + const std::vector empty_counts; + const std::vector wrong_count{4, 4, 4}; + const std::vector zero_count{12, 0}; + const std::vector negative_count{13, -1}; + const std::vector valid_counts{7, 5}; + + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, empty_counts), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, wrong_count), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, zero_count), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, negative_count), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(0, 2, valid_counts), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 0, valid_counts), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, valid_counts, 0), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, valid_counts, 2), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(12, 2, valid_counts, std::numeric_limits::max()), + std::invalid_argument); + + try { + const std::vector counts{4, 7, 8}; + (void)PipelineLayout::BuildCustomLayout(20, 3, counts); + FAIL() << "Expected std::invalid_argument"; + } catch (const std::invalid_argument &error) { + const std::string message = error.what(); + EXPECT_NE(message.find("received=[4,7,8]"), std::string::npos); + EXPECT_NE(message.find("num_stages=3"), std::string::npos); + EXPECT_NE(message.find("expected_sum=20"), std::string::npos); + EXPECT_NE(message.find("actual_sum=19"), std::string::npos); + } + + const std::vector overflowing_counts{std::numeric_limits::max(), 1}; + EXPECT_THROW((void)PipelineLayout::BuildCustomLayout(std::numeric_limits::max(), 2, overflowing_counts), + std::invalid_argument); +} + +TEST(PipelineLayoutTest, MapsGlobalLayersToStagesAndLocalIndices) { + const std::vector counts{7, 5}; + const PipelineLayout layout = PipelineLayout::BuildCustomLayout(12, 2, counts); + + EXPECT_EQ(layout.GetStageForLayer(0), 0); + EXPECT_EQ(layout.GetStageForLayer(6), 0); + EXPECT_EQ(layout.GetStageForLayer(7), 1); + EXPECT_EQ(layout.GetStageForLayer(11), 1); + EXPECT_EQ(layout.GetLocalLayerIndex(0, 0), 0); + EXPECT_EQ(layout.GetLocalLayerIndex(0, 6), 6); + EXPECT_EQ(layout.GetLocalLayerIndex(1, 7), 0); + EXPECT_EQ(layout.GetLocalLayerIndex(1, 11), 4); +} + +TEST(PipelineLayoutInvalidTest, RejectsOutOfRangeQueries) { + const std::vector counts{7, 5}; + const PipelineLayout layout = PipelineLayout::BuildCustomLayout(12, 2, counts); + + EXPECT_THROW((void)layout.GetStage(-1), std::out_of_range); + EXPECT_THROW((void)layout.GetStage(2), std::out_of_range); + EXPECT_THROW((void)layout.GetChunk(-1), std::out_of_range); + EXPECT_THROW((void)layout.GetChunk(layout.GetNumChunks()), std::out_of_range); + EXPECT_THROW((void)layout.GetStageForLayer(-1), std::out_of_range); + EXPECT_THROW((void)layout.GetStageForLayer(12), std::out_of_range); + EXPECT_THROW((void)layout.GetLocalLayerIndex(0, -1), std::out_of_range); + EXPECT_THROW((void)layout.GetLocalLayerIndex(0, 12), std::out_of_range); + EXPECT_THROW((void)layout.GetLocalLayerIndex(0, 7), std::out_of_range); +} + +// Explicit expected ranges/owners are independent of the factory implementation. +// Check every layer, including the transition into each stage's later chunks. +void ExpectChunks(const PipelineLayout &layout, const std::vector &bounds, const std::vector &owners) { + ASSERT_EQ(layout.GetNumStages(), 2); + ASSERT_EQ(bounds.size(), owners.size() + 1); + ASSERT_EQ(layout.GetNumChunks(), static_cast(owners.size())); + std::vector local_chunks(layout.GetNumStages(), 0); + std::vector local_layers(layout.GetNumStages(), 0); + for (int gid = 0; gid < layout.GetNumChunks(); ++gid) { + SCOPED_TRACE(gid); + const int owner = owners[gid]; + const auto &chunk = layout.GetChunk(gid); + EXPECT_EQ(chunk.global_chunk_id, gid); + EXPECT_EQ(chunk.stage_id, owner); + EXPECT_EQ(chunk.local_chunk_idx, local_chunks[owner]); + ASSERT_GT(layout.GetStage(owner).chunks.size(), static_cast(local_chunks[owner])); + EXPECT_EQ(&chunk, &layout.GetStage(owner).chunks[local_chunks[owner]++]); + ExpectRange(chunk.layer_range, bounds[gid], bounds[gid + 1]); + for (LayerIndex layer = bounds[gid]; layer < bounds[gid + 1]; ++layer) { + EXPECT_EQ(layout.GetStageForLayer(layer), owner); + EXPECT_EQ(layout.GetLocalLayerIndex(owner, layer), local_layers[owner]++); + EXPECT_THROW((void)layout.GetLocalLayerIndex(1 - owner, layer), std::out_of_range); + } + } + EXPECT_EQ(bounds.back(), layout.GetNumLayers()); + for (int stage = 0; stage < layout.GetNumStages(); ++stage) { + EXPECT_EQ(layout.GetStage(stage).chunks.size(), static_cast(local_chunks[stage])); + } +} + +TEST(PipelineLayoutTest, InterleavesUniformAndCustomChunks) { + const auto uniform = PipelineLayout::BuildUniformLayout(13, 2, 2); + ExpectChunks(uniform, {0, 4, 7, 10, 13}, {0, 1, 0, 1}); + const std::vector counts{2, 4, 3, 3}; + const auto custom = PipelineLayout::BuildCustomLayout(12, 2, counts, 2); + ExpectChunks(custom, {0, 2, 6, 9, 12}, {0, 1, 0, 1}); + for (const auto *layout : {&uniform, &custom}) { + EXPECT_EQ(layout->GetMaxLocalChunks(), 2); + EXPECT_EQ(layout->GetInputStage(), 0); + EXPECT_EQ(layout->GetOutputStage(), 1); + } +} + +TEST(PipelineLayoutTest, SupportsExplicitOwnersAndUnequalLocalChunkCounts) { + // Both endpoints move to stage 1; consecutive chunks and revisiting an owner + // must work without assuming round-robin placement or equal local counts. + const std::vector specs{{1, 2}, {0, 3}, {0, 4}, {1, 3}, {1, 1}}; + const auto layout = PipelineLayout::BuildChunkLayout(13, 2, specs); + ExpectChunks(layout, {0, 2, 5, 9, 12, 13}, {1, 0, 0, 1, 1}); + EXPECT_TRUE(layout.IsCustom()); + EXPECT_EQ(layout.GetMaxLocalChunks(), 3); + EXPECT_EQ(layout.GetInputStage(), 1); + EXPECT_EQ(layout.GetOutputStage(), 1); + for (int stage = 0; stage < 2; ++stage) { + EXPECT_EQ(layout.GetStage(stage).has_embedding, stage == 1); + EXPECT_EQ(layout.GetStage(stage).has_final_norm, stage == 1); + EXPECT_EQ(layout.GetStage(stage).has_lm_head, stage == 1); + } +} + +TEST(PipelineLayoutTest, PreservesEmptyAndUnderfilledDefaultMetadata) { + const auto empty = PipelineLayout::BuildUniformLayout(0, 2, 2); + EXPECT_EQ(empty.GetNumChunks(), 0); + EXPECT_EQ(empty.GetMaxLocalChunks(), 0); + EXPECT_EQ(empty.GetInputStage(), 0); + EXPECT_EQ(empty.GetOutputStage(), 1); + EXPECT_THROW((void)empty.GetChunk(0), std::out_of_range); + EXPECT_THROW((void)empty.GetStageForLayer(0), std::out_of_range); + EXPECT_THROW((void)empty.GetLocalLayerIndex(0, 0), std::out_of_range); + + const auto partial = PipelineLayout::BuildUniformLayout(1, 2, 2); + EXPECT_EQ(partial.GetNumChunks(), 1); + EXPECT_EQ(partial.GetMaxLocalChunks(), 1); + EXPECT_TRUE(partial.GetStage(1).chunks.empty()); + ExpectRange(partial.GetChunk(0).layer_range, 0, 1); + EXPECT_EQ(partial.GetStageForLayer(0), 0); + EXPECT_EQ(partial.GetLocalLayerIndex(0, 0), 0); + EXPECT_THROW((void)partial.GetLocalLayerIndex(1, 0), std::out_of_range); + EXPECT_THROW((void)partial.GetChunk(1), std::out_of_range); + EXPECT_EQ(partial.GetInputStage(), 0); + EXPECT_EQ(partial.GetOutputStage(), 1); // Legacy metadata, not empty-stage training. +} + +TEST(PipelineLayoutInvalidTest, RejectsInvalidExplicitLayouts) { + struct Case { + const char *reason; + std::vector specs; + }; + const std::vector cases{ + {"empty list", {}}, + {"negative owner", {{-1, 6}, {1, 6}}}, + {"owner out of range", {{0, 6}, {2, 6}}}, + {"zero layer count", {{0, 0}, {1, 12}}}, + {"negative layer count", {{0, -1}, {1, 13}}}, + {"sum too small", {{0, 5}, {1, 6}}}, + {"sum too large", {{0, 7}, {1, 6}}}, + {"empty stage", {{0, 6}, {0, 6}}}, + }; + for (const auto &test : cases) { + SCOPED_TRACE(test.reason); + EXPECT_THROW((void)PipelineLayout::BuildChunkLayout(12, 2, test.specs), std::invalid_argument); + } + const auto max = std::numeric_limits::max(); + const std::vector overflow{{0, max}, {1, 1}}; + EXPECT_THROW((void)PipelineLayout::BuildChunkLayout(max, 2, overflow), std::invalid_argument); + const std::vector valid{{0, 6}, {1, 6}}; + EXPECT_THROW((void)PipelineLayout::BuildChunkLayout(0, 2, valid), std::invalid_argument); + EXPECT_THROW((void)PipelineLayout::BuildChunkLayout(12, 0, valid), std::invalid_argument); +} + +} // namespace +} // namespace infini_train::nn::parallel diff --git a/tests/distributed/test_pipeline_loader.cc b/tests/distributed/test_pipeline_loader.cc new file mode 100644 index 000000000..4ebaf554f --- /dev/null +++ b/tests/distributed/test_pipeline_loader.cc @@ -0,0 +1,241 @@ +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#ifdef USE_OMP +#include +#endif + +#include "example/common/parser.h" +#include "example/gpt2/checkpoint_loader.h" +#include "example/llama3/checkpoint_loader.h" +#include "infini_train/include/nn/modules/transformer/transformer.h" +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::examples { +namespace { +namespace pp = nn::parallel; + +// Independent tiny LLMC fixture: each file element has a distinct, exactly +// representable value. Expected parameter values are recorded as it is written. +class CheckpointFixture { +public: + explicit CheckpointFixture(bool llama) { + path = std::filesystem::temp_directory_path() + / ("infinitrain-loader-" + std::to_string(std::chrono::steady_clock::now().time_since_epoch().count()) + + ".bin"); + std::ofstream out(path, std::ios::binary); + std::array header{}; + auto put = [&](int offset, auto value) { std::memcpy(header.data() + offset, &value, sizeof(value)); }; + put(0, llama ? 20240803 : 20240326); + put(4, 3); + put(8, 4); + put(12, 16); + put(16, 4); + put(20, 2); + if (llama) { + put(24, 1); + put(28, 8); // GQA: 2 Q heads, 1 KV head. + put(32, 1.0f); + put(36, 8); + put(40, 1e-5f); + put(44, 10000.0f); + put(48, 0); + put(52, 2); + put(56, 3); + put(60, 0); + } else { + put(24, 8); + put(28, 20); // Include padded WTE rows to check file positioning. + } + out.write(header.data(), header.size()); + float next = 1; + auto append = [&](const std::string &name, int count) { + std::vector values(count); + for (auto &v : values) { v = next++; } + out.write(reinterpret_cast(values.data()), count * sizeof(float)); + if (!name.empty()) { + expected[name] = std::move(values); + } + }; + append("transformer.wte.weight", 16 * 8); + if (!llama) { + expected["lm_head.weight"] = expected.at("transformer.wte.weight"); + append("", 4 * 8); + append("transformer.wpe.weight", 4 * 8); + } + auto layers = [&](const std::string &suffix, int count) { + for (int layer = 0; layer < 4; ++layer) { + append("transformer.h." + std::to_string(layer) + "." + suffix, count); + } + }; + layers("ln_1.weight", 8); + if (!llama) { + layers("ln_1.bias", 8); + } + layers("attn.c_attn.weight", (llama ? 16 : 24) * 8); + if (!llama) { + layers("attn.c_attn.bias", 24); + } + layers("attn.c_proj.weight", 8 * 8); + if (!llama) { + layers("attn.c_proj.bias", 8); + } + layers("ln_2.weight", 8); + if (!llama) { + layers("ln_2.bias", 8); + } + layers("mlp.c_fc.weight", (llama ? 24 : 32) * 8); + if (llama) { + layers("mlp.c_fc2.weight", 24 * 8); + } else { + layers("mlp.c_fc.bias", 32); + } + layers("mlp.c_proj.weight", 8 * (llama ? 24 : 32)); + if (!llama) { + layers("mlp.c_proj.bias", 8); + } + append("transformer.ln_f.weight", 8); + if (llama) { + append("lm_head.weight", 16 * 8); + } else { + append("transformer.ln_f.bias", 8); + } + if (!out) { + throw std::runtime_error("cannot write checkpoint fixture"); + } + } + ~CheckpointFixture() { + std::error_code error; + std::filesystem::remove(path, error); + } + std::filesystem::path path; + std::map> expected; +}; + +std::shared_ptr Load(bool llama, const CheckpointFixture &fixture, + const std::optional &request) { + if (llama) { + return request ? llama3::LoadFromLLMC(fixture.path.string(), *request) + : llama3::LoadFromLLMC(fixture.path.string()); + } + return request ? gpt2::LoadFromLLMC(fixture.path.string(), *request) : gpt2::LoadFromLLMC(fixture.path.string()); +} + +TEST(PipelineLoaderTest, LoadsExactOwnedValuesAcrossAllRequestSources) { + const int stages = pp::global::GetPipelineParallelSize(); + const int vpp = pp::global::GetVirtualPipelineParallelSize(); + std::vector> requests{std::nullopt, PipelineLayoutRequest{std::monostate{}}}; + if (stages == 1) { + requests.emplace_back(ParsePipelineLayoutRequest("4", "", stages, vpp)); + } else if (vpp == 1) { + requests.emplace_back(ParsePipelineLayoutRequest("1,3", "", stages, vpp)); + requests.emplace_back(ParsePipelineLayoutRequest("", "1:1,0:3", stages, vpp)); + } else { + requests.emplace_back(ParsePipelineLayoutRequest("1,1,1,1", "", stages, vpp)); + requests.emplace_back(ParsePipelineLayoutRequest("", "1:1,0:1,1:2", stages, vpp)); + } + for (bool llama : {false, true}) { + CheckpointFixture fixture(llama); + for (std::size_t mode = 0; mode < requests.size(); ++mode) { + std::size_t checked = 0; + for (int stage = 0; stage < stages; ++stage) { + SCOPED_TRACE(::testing::Message() << "llama=" << llama << " mode=" << mode << " stage=" << stage); + pp::pp_rank = stage; + const auto model = Load(llama, fixture, requests[mode]); + auto layout = model->GetPipelineLayout(); + EXPECT_EQ(static_cast(layout), requests[mode].has_value()); + if (!layout) { + layout = std::make_shared( + pp::PipelineLayout::BuildUniformLayout(4, stages, vpp)); + } + EXPECT_EQ(model->Config().n_layer, 4); + EXPECT_EQ(model->Config().n_embd, 8); + EXPECT_EQ(model->Config().original_vocab_size, 16); + if (!llama) { + EXPECT_EQ(model->Config().tie_weights, stages == 1); + } + const auto &stage_layout = layout->GetStage(stage); + std::vector globals; + for (const auto &chunk : stage_layout.chunks) { + for (auto layer = chunk.layer_range.begin; layer < chunk.layer_range.end; ++layer) { + globals.push_back(layer); + } + } + const auto state = model->StateDict(); + ASSERT_FALSE(state.empty()); + if (!llama && stage_layout.has_embedding && stage_layout.has_lm_head) { + const auto &wte = state.at("transformer.wte.weight"); + const auto &head = state.at("lm_head.weight"); + EXPECT_EQ(wte.get() == head.get(), stages == 1); + } + for (const auto &[local_name, tensor] : state) { + std::string name = local_name; + const std::string prefix = "transformer.h."; + if (name.starts_with(prefix)) { + const auto dot = name.find('.', prefix.size()); + const int local = std::stoi(name.substr(prefix.size(), dot - prefix.size())); + name = prefix + std::to_string(globals.at(local)) + name.substr(dot); + } + // The causal mask is a constructed buffer, not checkpoint data. + if (name.ends_with(".attn.bias")) { + ASSERT_EQ(tensor->Dims(), (std::vector{1, 1, 4, 4})); + const auto *mask = static_cast(tensor->DataPtr()); + for (int row = 0; row < 4; ++row) { + for (int col = 0; col < 4; ++col) { + ASSERT_EQ(mask[row * 4 + col], col <= row ? 1.0f : 0.0f); + } + } + continue; + } + ASSERT_TRUE(fixture.expected.contains(name)) << name; + const auto &expected = fixture.expected.at(name); + ASSERT_EQ(tensor->NumElements(), expected.size()) << name; + const auto *actual = static_cast(tensor->DataPtr()); + for (std::size_t i = 0; i < expected.size(); ++i) { + ASSERT_EQ(actual[i], expected[i]) << name << " element=" << i; + } + ++checked; + } + } + EXPECT_EQ(checked, fixture.expected.size()); // No missing or extra parameters across stages. + } + } + pp::pp_rank = 0; +} + +TEST(PipelineLoaderTest, RejectsLayerSumFromCheckpointHeader) { + pp::pp_rank = 0; + const int stages = pp::global::GetPipelineParallelSize(); + const int vpp = pp::global::GetVirtualPipelineParallelSize(); + std::vector counts(stages * vpp, 1); + counts.front() += 4; // Syntactically legal, but sum disagrees with the 4-layer file. + for (bool llama : {false, true}) { + CheckpointFixture fixture(llama); + EXPECT_THROW((void)Load(llama, fixture, PipelineLayoutRequest{counts}), std::invalid_argument); + } +} +} // namespace +} // namespace infini_train::examples + +// Each CTest process has a separate GlobalEnv; no communicator or GPU is created. +int main(int argc, char **argv) { + ::testing::InitGoogleTest(&argc, argv); +#ifdef USE_OMP + omp_set_num_threads(1); +#endif + const int pp = (std::getenv("PIPELINE_LOADER_TEST_PP") ? std::atoi(std::getenv("PIPELINE_LOADER_TEST_PP")) : 1); + const int vpp = (std::getenv("PIPELINE_LOADER_TEST_VPP") ? std::atoi(std::getenv("PIPELINE_LOADER_TEST_VPP")) : 1); + infini_train::nn::parallel::global::InitAllEnv(1, 1, false, pp, vpp); + return RUN_ALL_TESTS(); +} diff --git a/tests/distributed/test_pipeline_parallel.cc b/tests/distributed/test_pipeline_parallel.cc new file mode 100644 index 000000000..64ea83796 --- /dev/null +++ b/tests/distributed/test_pipeline_parallel.cc @@ -0,0 +1,221 @@ +#include "gtest/gtest.h" + +#include +#include +#include +#include + +#ifdef USE_OMP +#include +#endif + +#include "infini_train/include/nn/modules/container.h" +#include "infini_train/include/nn/modules/transformer/transformer.h" +#include "infini_train/include/nn/parallel/pp/pipeline_parallel.h" +#include "infini_train/include/nn/parallel/pp/pipeline_schedule.h" +#include "infini_train/include/nn/parallel/pp/pipeline_stage.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::nn::parallel { +namespace { + +TransformerConfig SmallConfig(bool llama = false) { + TransformerConfig config; + config.n_layer = 6; + config.n_embd = 8; + config.n_head = 2; + config.n_kv_head = 2; + config.vocab_size = config.original_vocab_size = 16; + config.block_size = 8; + config.multiple_of = 8; + if (llama) { + config.position_embedding_type = PositionEmbeddingType::kRoPE; + config.norm_type = NormType::kRMSNorm; + config.activation_type = MLPType::kSwiGLU; + config.add_bias_linear = false; + config.tie_weights = false; + } + return config; +} + +class PipelineParallelTest : public ::testing::Test { +protected: + void SetUp() override { +#ifdef USE_OMP + threads_ = omp_get_max_threads(); + omp_set_num_threads(1); +#endif + } + void TearDown() override { +#ifdef USE_OMP + omp_set_num_threads(threads_); +#endif + } + const std::vector> shapes_{{2, 4, 8}}; +#ifdef USE_OMP + int threads_; +#endif +}; + +// Verify the real model's layers and the wrapper's pointer identity/order. No GPU +// communication is needed to catch wrong stage initialization or copied modules. +TEST_F(PipelineParallelTest, BuildsAndWrapsActualModelChunks) { + struct Case { + const char *name; + PipelineLayout layout; + bool explicit_mode; + std::vector> chunk_sizes; + int input_stage; + int output_stage; + }; + const std::vector two_stage_counts{1, 2, 1, 2}; + const std::vector three_stage_counts{1, 2, 3}; + const std::vector three_stage_vpp_counts{1, 2, 1, 1, 1, 3}; + const std::vector two_stage_specs{{1, 1}, {0, 2}, {1, 1}, {1, 2}}; + const std::vector three_stage_specs{{1, 1}, {0, 1}, {2, 1}, {1, 1}, {2, 2}}; + // Literal expectations keep the model/wrapper assertions independent of layout queries. + const std::vector cases{ + {"PP2 uniform", PipelineLayout::BuildUniformLayout(6, 2), false, {{3}, {3}}, 0, 1}, + {"PP2 custom vPP2", PipelineLayout::BuildCustomLayout(6, 2, two_stage_counts, 2), + false, {{1, 1}, {2, 2}}, 0, 1}, + {"PP2 explicit", PipelineLayout::BuildChunkLayout(6, 2, two_stage_specs), + true, {{2}, {1, 1, 2}}, 1, 1}, + {"PP3 uniform", PipelineLayout::BuildUniformLayout(8, 3), false, {{3}, {3}, {2}}, 0, 2}, + {"PP3 custom", PipelineLayout::BuildCustomLayout(6, 3, three_stage_counts), + false, {{1}, {2}, {3}}, 0, 2}, + {"PP3 uniform vPP2", PipelineLayout::BuildUniformLayout(9, 3, 2), + false, {{2, 1}, {2, 1}, {2, 1}}, 0, 2}, + {"PP3 custom vPP2", PipelineLayout::BuildCustomLayout(9, 3, three_stage_vpp_counts, 2), + false, {{1, 1}, {2, 1}, {1, 3}}, 0, 2}, + {"PP3 explicit", PipelineLayout::BuildChunkLayout(6, 3, three_stage_specs), + true, {{1}, {1, 1}, {1, 2}}, 1, 2}, + }; + for (bool llama : {false, true}) { + for (const auto &test : cases) { + const auto layout = std::make_shared(test.layout); + const int stages = static_cast(test.chunk_sizes.size()); + auto config = SmallConfig(llama); + config.n_layer = static_cast(layout->GetNumLayers()); + for (int stage_id = 0; stage_id < stages; ++stage_id) { + SCOPED_TRACE(::testing::Message() << "llama=" << llama << " case=" << test.name << " stage=" << stage_id); + auto model = std::make_shared(config, layout, stage_id); + EXPECT_EQ(model->GetStageId(), stage_id); + EXPECT_EQ(model->GetPipelineLayout(), layout); + PipelineParallel wrapper(model, stages, 2, shapes_, stage_id, Device(), layout, test.explicit_mode); + EXPECT_EQ(wrapper.GetPipelineLayout(), layout); // Regression for moved-from pointer check. + const auto &stage = layout->GetStage(stage_id); + const auto &chunks = *wrapper.mutable_chunks(); + ASSERT_EQ(chunks.size(), test.chunk_sizes[stage_id].size()); + EXPECT_EQ(stage.has_embedding, stage_id == test.input_stage); + EXPECT_EQ(stage.has_final_norm, stage_id == test.output_stage); + EXPECT_EQ(stage.has_lm_head, stage_id == test.output_stage); + if (stage_id != test.input_stage) { + EXPECT_THROW((void)model->mutable_module(Module::kPPFirstStageName), std::out_of_range); + } + if (stage_id != test.output_stage) { + EXPECT_THROW((void)model->mutable_module(Module::kPPLastStageName), std::out_of_range); + } + EXPECT_EQ(model->GetNumChunks(), static_cast(chunks.size())); + auto all_layers = std::dynamic_pointer_cast( + model->mutable_module(TransformerModel::kTransformerModelName) + ->mutable_module(TransformerChunk::kHLayerName)); + ASSERT_NE(all_layers, nullptr); + std::size_t offset = 0; + for (std::size_t local = 0; local < chunks.size(); ++local) { + auto original = model->mutable_module(Module::kPPChunkNamePrefix + std::to_string(local)); + auto layers = std::dynamic_pointer_cast( + original->mutable_module(TransformerChunk::kHLayerName)); + ASSERT_NE(layers, nullptr); + ASSERT_EQ(std::distance(layers->begin(), layers->end()), test.chunk_sizes[stage_id][local]); + for (const auto &layer : *layers) { EXPECT_EQ(layer, (*all_layers)[offset++]); } + std::vector> expected; + if (local == 0 && stage.has_embedding) { + expected.push_back(model->mutable_module(Module::kPPFirstStageName)); + } + expected.push_back(original); + if (local + 1 == chunks.size() && stage.has_lm_head) { + expected.push_back(model->mutable_module(Module::kPPLastStageName)); + } + EXPECT_EQ(chunks[local]->type(), Sequential::kType); + for (std::size_t part = 0; part < expected.size(); ++part) { + EXPECT_EQ(chunks[local]->mutable_module(std::to_string(part)), expected[part]); + } + EXPECT_THROW((void)chunks[local]->mutable_module(std::to_string(expected.size())), + std::out_of_range); + } + EXPECT_EQ(std::distance(all_layers->begin(), all_layers->end()), offset); + } + } + } +} + +TEST_F(PipelineParallelTest, SingleStageRetainsLegacyModelAndWrapper) { + auto config = SmallConfig(); + auto legacy_model = std::make_shared(config); + auto layout = std::make_shared(PipelineLayout::BuildUniformLayout(6, 1)); + auto model = std::make_shared(config, layout, 0); + EXPECT_EQ(legacy_model->GetStageId(), 0); + const auto old_state = legacy_model->StateDict(); + const auto new_state = model->StateDict(); + ASSERT_FALSE(old_state.empty()); + ASSERT_EQ(old_state.size(), new_state.size()); + for (const auto &[name, parameter] : old_state) { + ASSERT_TRUE(new_state.contains(name)) << name; + EXPECT_EQ(parameter->Dims(), new_state.at(name)->Dims()); + } + PipelineParallel legacy(legacy_model, 1, 2, shapes_, 0, Device(), 1); + PipelineParallel current(model, 1, 2, shapes_, 0, Device(), layout); + ASSERT_EQ(legacy.mutable_chunks()->size(), 1); + ASSERT_EQ(current.mutable_chunks()->size(), 1); + for (int part = 0; part < 3; ++part) { + EXPECT_EQ(legacy.mutable_chunks()->at(0)->module(std::to_string(part)).type(), + current.mutable_chunks()->at(0)->module(std::to_string(part)).type()); + } +} + +TEST_F(PipelineParallelTest, RejectsInvalidModelAndWrapperInputs) { + auto layout = std::make_shared(PipelineLayout::BuildUniformLayout(6, 2)); + const auto config = SmallConfig(); + EXPECT_THROW((void)TransformerModel(config, nullptr, 0), std::invalid_argument); + EXPECT_THROW((void)TransformerModel(config, layout, 2), std::out_of_range); + auto wrong_config = config; + wrong_config.n_layer = 7; + EXPECT_THROW((void)TransformerModel(wrong_config, layout, 0), std::invalid_argument); + auto model = std::make_shared(config, layout, 0); + EXPECT_THROW((void)PipelineParallel(model, 2, 2, shapes_, 0, Device(), nullptr), std::invalid_argument); + EXPECT_THROW((void)PipelineParallel(nullptr, 2, 2, shapes_, 0, Device(), layout), std::invalid_argument); + EXPECT_THROW((void)PipelineParallel(model, 3, 2, shapes_, 0, Device(), layout), std::invalid_argument); + EXPECT_THROW((void)PipelineParallel(model, 2, 2, shapes_, 2, Device(), layout), std::out_of_range); + EXPECT_THROW((void)PipelineParallel(model, 2, 0, shapes_, 0, Device(), layout), std::invalid_argument); + const std::vector other_counts{1, 5}; + auto other_layout = std::make_shared(PipelineLayout::BuildCustomLayout(6, 2, other_counts)); + EXPECT_THROW((void)PipelineParallel(model, 2, 2, shapes_, 0, Device(), other_layout), std::invalid_argument); + EXPECT_THROW((void)PipelineParallel(model, 2, 2, shapes_, 1, Device(), layout), std::invalid_argument); +} + +TEST_F(PipelineParallelTest, LegacyStageInfoPreservesInterleavedRanges) { + const auto stage0 = PipelineParallel::GetStageInfo(13, 2, 0, 2); + const auto stage1 = PipelineParallel::GetStageInfo(13, 2, 1, 2); + EXPECT_EQ(stage0.layer_ranges_per_chunk, (std::vector>{{0, 4}, {7, 10}})); + EXPECT_EQ(stage1.layer_ranges_per_chunk, (std::vector>{{4, 7}, {10, 13}})); + EXPECT_TRUE(stage0.is_first_stage); + EXPECT_FALSE(stage0.is_last_stage); + EXPECT_FALSE(stage1.is_first_stage); + EXPECT_TRUE(stage1.is_last_stage); + EXPECT_TRUE(PipelineParallel::GetStageInfo(0, 2, 0, 2).layer_ranges_per_chunk.empty()); +} + +// Invalid requests must fail before any point-to-point communication. +TEST_F(PipelineParallelTest, ScheduleRejectsInconsistentModeAndInvalidInputs) { + const std::vector specs{{1, 2}, {0, 2}, {1, 2}}; + auto layout = std::make_shared(PipelineLayout::BuildChunkLayout(6, 2, specs)); + std::vector> chunks{std::make_shared()}; + auto stage = std::make_shared(0, 2, shapes_, Device(), std::move(chunks)); + EXPECT_THROW((void)PipelineSchedule(stage, 2, 2, layout, false), std::invalid_argument); + EXPECT_THROW((void)PipelineSchedule(stage, 2, 2, nullptr, true), std::invalid_argument); + PipelineSchedule schedule(stage, 2, 2, layout, true); + EXPECT_THROW((void)schedule.StepMicroBatches({}, {}, nullptr, DataType::kFLOAT32), std::invalid_argument); +} + +} // namespace +} // namespace infini_train::nn::parallel diff --git a/tests/distributed/test_pipeline_schedule.cc b/tests/distributed/test_pipeline_schedule.cc new file mode 100644 index 000000000..be1e93f1a --- /dev/null +++ b/tests/distributed/test_pipeline_schedule.cc @@ -0,0 +1,389 @@ +#include "gtest/gtest.h" +#include +#include +#include + +#include +#include +#include +#include +#include + +#include "infini_train/include/autograd/function_hook.h" +#include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/pp/pipeline_layout.h" +#include "infini_train/include/nn/parallel/pp/pipeline_schedule.h" +#include "infini_train/include/nn/parallel/pp/pipeline_stage.h" +#include "infini_train/include/optimizer.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::nn::parallel { +namespace { +using Scheduler = PipelineParallelScheduler; + +// Model ProcessGroup's stream-ordered, unbuffered P2P operations. Send/Recv +// enqueue a compute-stream wait, so a later operation cannot unblock an earlier +// unmatched operation on that same rank. This is stronger than tensor DAG order. +std::string PendingRendezvous(const std::vector &tasks, int stages, const std::vector &owners) { + struct Event { + bool send; + int peer; + int mb; + int edge; + bool forward; + }; + std::vector> queues(stages); + const int chunks = static_cast(owners.size()); + for (const auto &task : tasks) { + const int gid = task.global_chunk_id, rank = owners.at(gid), mb = task.microbatch_id; + auto &events = queues.at(rank); + if (task.is_forward) { + if (gid > 0 && owners[gid - 1] != rank) { + events.push_back({false, owners[gid - 1], mb, gid - 1, true}); + } + if (gid + 1 < chunks && owners[gid + 1] != rank) { + events.push_back({true, owners[gid + 1], mb, gid, true}); + } + } else { + if (gid + 1 < chunks && owners[gid + 1] == rank) { + continue; + } + if (gid + 1 < chunks) { + events.push_back({false, owners[gid + 1], mb, gid, false}); + } + int first = gid; + while (first > 0 && owners[first - 1] == rank) { --first; } + if (first > 0) { + events.push_back({true, owners[first - 1], mb, first - 1, false}); + } + } + } + for (;;) { + bool progress = false; + for (int rank = 0; rank < stages; ++rank) { + if (queues[rank].empty()) { + continue; + } + const auto a = queues[rank].front(); + if (queues[a.peer].empty()) { + continue; + } + const auto b = queues[a.peer].front(); + if (a.send != b.send && b.peer == rank && a.mb == b.mb && a.edge == b.edge && a.forward == b.forward) { + queues[rank].pop_front(); + queues[a.peer].pop_front(); + progress = true; + } + } + if (progress) { + continue; + } + std::ostringstream pending; + for (int rank = 0; rank < stages; ++rank) { + if (queues[rank].empty()) { + continue; + } + const auto &e = queues[rank].front(); + pending << "rank=" << rank << (e.send ? " send to=" : " recv from=") << e.peer << " mb=" << e.mb + << " edge=" << e.edge << (e.forward ? " F; " : " B; "); + } + return pending.str(); + } +} + +TEST(PipelineScheduleTest, CommunicationOrderMakesProgressAcrossAllOwners) { + // Exhaust two-stage owners through six chunks and three-stage owners through five. + for (int stages : {2, 3}) { + for (int count = stages; count <= (stages == 2 ? 6 : 5); ++count) { + int combinations = 1; + for (int i = 0; i < count; ++i) { combinations *= stages; } + for (int code = 0; code < combinations; ++code) { + std::vector owners, present(stages, 0); + std::vector specs; + int digits = code; + for (int gid = 0; gid < count; ++gid) { + const int owner = digits % stages; + digits /= stages; + owners.push_back(owner); + ++present[owner]; + specs.push_back({owner, 1}); + } + if (std::find(present.begin(), present.end(), 0) != present.end()) { + continue; + } + const auto layout = PipelineLayout::BuildChunkLayout(count, stages, specs); + for (int n = 1; n <= 4; ++n) { + SCOPED_TRACE(::testing::Message() << "stages=" << stages << " chunks=" << count + << " owners=" << code << " microbatches=" << n); + const auto tasks = Scheduler::GenerateGPipeSchedule(n, stages, layout.GetMaxLocalChunks(), &layout); + ASSERT_EQ(PendingRendezvous(tasks, stages, owners), ""); + } + } + } + } +} + +TEST(PipelineScheduleTest, DetectsLegacyCyclicCommunicationDeadlock) { + // Fixed old PP2/vPP2/n2 forward sequence; do not derive it from the new generator. + const std::vector> rows{{0, 0, 0}, {1, 0, 1}, {1, 1, 0}, {2, 1, 1}, + {2, 0, 2}, {3, 0, 3}, {3, 1, 2}, {4, 1, 3}}; + std::vector tasks; + for (const auto &row : rows) { tasks.push_back(Scheduler::CreateTask(row[0], row[1], row[2], 2, 4, true)); } + EXPECT_EQ(PendingRendezvous(tasks, 2, {0, 1, 0, 1}), + "rank=0 send to=1 mb=1 edge=0 F; rank=1 send to=0 mb=0 edge=1 F; "); +} + +TEST(PipelineScheduleTest, CyclicGPipeUsesProgressSafeOrderAcrossAllRequestSources) { + const auto uniform = PipelineLayout::BuildUniformLayout(12, 2, 2); + const std::vector counts{1, 2, 3, 6}; + const auto custom = PipelineLayout::BuildCustomLayout(12, 2, counts, 2); + const std::vector specs{{0, 1}, {1, 2}, {0, 3}, {1, 6}}; + const auto explicit_layout = PipelineLayout::BuildChunkLayout(12, 2, specs); + const std::vector> expected{ + {0, 0, 0}, {1, 0, 1}, {2, 0, 2}, {3, 0, 3}, {4, 1, 0}, {5, 1, 1}, {6, 1, 2}, {7, 1, 3}, + {8, 1, 3}, {9, 1, 2}, {10, 1, 1}, {11, 1, 0}, {12, 0, 3}, {13, 0, 2}, {14, 0, 1}, {15, 0, 0}}; + for (const auto *layout : {static_cast(nullptr), &uniform, &custom, &explicit_layout}) { + const auto tasks = Scheduler::GenerateGPipeSchedule(2, 2, 2, layout); + ASSERT_EQ(tasks.size(), expected.size()); + for (size_t i = 0; i < tasks.size(); ++i) { + const auto &task = tasks[i]; + EXPECT_EQ((std::array{task.step, task.microbatch_id, task.global_chunk_id}), expected[i]); + EXPECT_EQ(task.is_forward, i < 8); + EXPECT_EQ(task.stage_id, expected[i][2] % 2); + EXPECT_EQ(task.local_chunk_idx, expected[i][2] / 2); + } + EXPECT_EQ(PendingRendezvous(tasks, 2, {0, 1, 0, 1}), ""); + } +} + +TEST(PipelineScheduleTest, PreservesPinnedLegacyTaskOrder) { + // Captured from f604940's generators, not from the modified implementation. + // Each row is (step, microbatch, global chunk), with forward flags listed separately. + struct Case { + int vpp; + bool one_f_one_b; + std::vector> rows; + const char *directions; + }; + const std::vector cases{ + {1, + false, + {{0, 0, 0}, {1, 0, 1}, {1, 1, 0}, {2, 1, 1}, {3, 1, 1}, {4, 0, 1}, {4, 1, 0}, {5, 0, 0}}, + "FFFFBBBB"}, + {1, true, {{0, 0, 0}, {1, 0, 1}, {1, 0, 1}, {1, 1, 0}, {2, 0, 0}, {2, 1, 1}, {2, 1, 1}, {3, 1, 0}}, "FFBFBFBB"}, + {2, + true, + {{0, 0, 0}, + {1, 0, 1}, + {1, 1, 0}, + {2, 0, 2}, + {2, 1, 1}, + {3, 0, 3}, + {3, 0, 3}, + {3, 1, 2}, + {4, 0, 2}, + {4, 1, 3}, + {4, 1, 3}, + {5, 0, 1}, + {5, 1, 2}, + {6, 0, 0}, + {6, 1, 1}, + {7, 1, 0}}, + "FFFFFFBFBFBBBBBB"}, + }; + for (const auto &c : cases) { + const auto uniform = PipelineLayout::BuildUniformLayout(12, 2, c.vpp); + const std::vector counts + = c.vpp == 1 ? std::vector{3, 9} : std::vector{1, 2, 3, 6}; + const auto custom = PipelineLayout::BuildCustomLayout(12, 2, counts, c.vpp); + std::vector specs; + for (size_t gid = 0; gid < counts.size(); ++gid) { specs.push_back({static_cast(gid % 2), counts[gid]}); } + const auto explicit_interleaved = PipelineLayout::BuildChunkLayout(12, 2, specs); + for (const auto *layout : + {static_cast(nullptr), &uniform, &custom, &explicit_interleaved}) { + SCOPED_TRACE(::testing::Message() << "vpp=" << c.vpp << " 1F1B=" << c.one_f_one_b << " layout=" << layout); + const auto tasks = c.one_f_one_b ? Scheduler::GenerateInterleaved1F1BSchedule(2, 2, c.vpp, layout) + : Scheduler::GenerateGPipeSchedule(2, 2, c.vpp, layout); + ASSERT_EQ(tasks.size(), c.rows.size()); + for (size_t i = 0; i < tasks.size(); ++i) { + const auto &t = tasks[i]; + EXPECT_EQ((std::array{t.step, t.microbatch_id, t.global_chunk_id}), c.rows[i]); + EXPECT_EQ(t.is_forward, c.directions[i] == 'F'); + EXPECT_EQ(t.stage_id, c.rows[i][2] % 2); + EXPECT_EQ(t.local_chunk_idx, c.rows[i][2] / 2); + EXPECT_EQ(t.is_first_chunk, c.rows[i][2] == 0); + EXPECT_EQ(t.is_last_chunk, c.rows[i][2] == 2 * c.vpp - 1); + } + } + } +} + +TEST(PipelineScheduleTest, ArbitraryMappingsPreserveDependenciesAndOwnership) { + // Local chains, revisited stages, moved endpoints and unequal local chunk counts. + const std::vector> mappings{{0, 1, 1, 0}, {1, 0, 1, 1}, {2, 0, 0, 1, 2}}; + for (const auto &owners : mappings) { + const int stages = owners == mappings.back() ? 3 : 2; + const int count = static_cast(owners.size()); + std::vector specs; + std::vector locals, next_local(stages, 0); + for (int owner : owners) { + specs.push_back({owner, 1}); + locals.push_back(next_local[owner]++); + } + const auto layout = PipelineLayout::BuildChunkLayout(count, stages, specs); + for (int n : {1, 3}) { + const auto tasks = Scheduler::GenerateGPipeSchedule(n, stages, layout.GetMaxLocalChunks(), &layout); + ASSERT_EQ(tasks.size(), 2 * n * count); + std::vector> forwards(n, std::vector(count, -1)); + auto backwards = forwards; + bool backward_started = false; + for (size_t i = 0; i < tasks.size(); ++i) { + const auto &t = tasks[i]; + ASSERT_GE(t.microbatch_id, 0); + ASSERT_LT(t.microbatch_id, n); + ASSERT_GE(t.global_chunk_id, 0); + ASSERT_LT(t.global_chunk_id, count); + const int gid = t.global_chunk_id, mb = t.microbatch_id; + EXPECT_EQ(t.stage_id, owners[gid]); + EXPECT_EQ(t.local_chunk_idx, locals[gid]); + EXPECT_EQ(t.is_first_chunk, gid == 0); + EXPECT_EQ(t.is_last_chunk, gid == count - 1); + auto &seen = t.is_forward ? forwards : backwards; + EXPECT_EQ(seen[mb][gid], -1); + seen[mb][gid] = static_cast(i); + if (t.is_forward) { + EXPECT_FALSE(backward_started); + if (gid > 0) { + EXPECT_GE(forwards[mb][gid - 1], 0); + } + } else { + backward_started = true; + EXPECT_GE(forwards[mb][gid], 0); + if (gid + 1 < count) { + EXPECT_GE(backwards[mb][gid + 1], 0); + } + } + } + } + } +} + +TEST(PipelineScheduleTest, RejectsMismatchedDimensionsAndOverflowBeforeAllocation) { + const auto layout = PipelineLayout::BuildUniformLayout(8, 2, 2); + EXPECT_THROW(Scheduler::GenerateGPipeSchedule(1, 3, 2, &layout), std::invalid_argument); + EXPECT_THROW(Scheduler::GenerateGPipeSchedule(1, 2, 1, &layout), std::invalid_argument); + EXPECT_THROW(Scheduler::GenerateGPipeSchedule(std::numeric_limits::max(), 2, 2), std::invalid_argument); + EXPECT_THROW(Scheduler::GenerateGPipeSchedule(std::numeric_limits::max() / 8 + 1, 2, 2), + std::invalid_argument); + EXPECT_TRUE(Scheduler::GenerateGPipeSchedule(0, 2, 2).empty()); + EXPECT_TRUE(Scheduler::GenerateInterleaved1F1BSchedule(0, 2, 2).empty()); +} + +struct SyncState { + bool suppressed = false; + int entries = 0; + int exits = 0; + std::vector suppressed_at_accumulation; + std::vector suppressed_at_backward; +}; + +class ObserveAccumulation : public autograd::PostAccumulateGradHook { +public: + explicit ObserveAccumulation(std::shared_ptr state) : state_(std::move(state)) {} + void operator()(const std::shared_ptr &) override { + state_->suppressed_at_accumulation.push_back(state_->suppressed); + } + +private: + std::shared_ptr state_; +}; + +// CPU numerical test of the real executor. No mocked communication or copied executor. +class Scale : public Module { +public: + explicit Scale(float value) { + parameters_["weight"] = std::make_shared(&value, std::vector{1}, DataType::kFLOAT32, Device()); + parameters_["weight"]->set_requires_grad(true); + parameters_["weight"]->RegisterPostAccumulateGradHook(std::make_shared(sync)); + backward_hook_ = RegisterBackwardPreHook( + [state = sync](Module *, const auto &) { state->suppressed_at_backward.push_back(state->suppressed); }); + } + std::unique_ptr no_sync() override { + sync->suppressed = true; + ++sync->entries; + return std::make_unique([state = sync] { + state->suppressed = false; + ++state->exits; + }); + } + std::shared_ptr sync = std::make_shared(); + std::shared_ptr backward_hook_; + std::vector> Forward(const std::vector> &x) override { + return {x[0] * parameters_["weight"]}; + } +}; +class SquaredLoss : public Module { +public: + std::vector> Forward(const std::vector> &x) override { + auto error = x[0] - x[1]; + return {(error * error)->Mean(0)}; + } +}; +TEST(PipelineScheduleTest, LocalChainAccumulatesGradientsAndUpdatesOncePerStep) { + auto first = std::make_shared(2.0f), second = std::make_shared(3.0f); + std::vector> chunks{first, second}; + auto stage + = std::make_shared(0, 1, std::vector>{{1}}, Device(), std::move(chunks)); + const std::vector specs{{0, 1}, {0, 1}}; + auto layout = std::make_shared(PipelineLayout::BuildChunkLayout(2, 1, specs)); + PipelineSchedule schedule(stage, 1, 2, layout, true); + auto p = first->parameter("weight"), q = second->parameter("weight"); + auto optimizer = std::make_shared(std::vector>{p, q}, 0.01f); + auto loss_fn = std::make_shared(); + const float xs[]{1, 2}, ys[]{0, 0}; + auto x = std::make_shared(xs, std::vector{2}, DataType::kFLOAT32, Device()); + auto y = std::make_shared(ys, std::vector{2}, DataType::kFLOAT32, Device()); + const auto scalar = [](const auto &t) { return static_cast(t->DataPtr())[0]; }; + for (int step = 0; step < 2; ++step) { + const float a = scalar(p), b = scalar(q); + // mean((a*b*[1,2])^2) = 2.5*a^2*b^2. + const float loss = schedule.Step(x, y, optimizer, loss_fn, DataType::kFLOAT32); + EXPECT_NEAR(loss, 2.5f * a * a * b * b, 1e-4f); + ASSERT_NE(p->grad(), nullptr); + ASSERT_NE(q->grad(), nullptr); + EXPECT_NEAR(scalar(p->grad()), 5 * a * b * b, 1e-4f); + EXPECT_NEAR(scalar(q->grad()), 5 * a * a * b, 1e-4f); + EXPECT_NEAR(scalar(p), a - 0.01f * 5 * a * b * b, 1e-5f); + EXPECT_NEAR(scalar(q), b - 0.01f * 5 * a * a * b, 1e-5f); + } + for (const auto &chunk : {first, second}) { + EXPECT_EQ(chunk->sync->entries, 2); + EXPECT_EQ(chunk->sync->exits, 2); + // Every chunk in the chain must leave no_sync before its final backward. + EXPECT_EQ(chunk->sync->suppressed_at_backward, (std::vector{true, false, true, false})); + // GPipe builds all forward graphs first; the shared accumulator fires once per step. + EXPECT_EQ(chunk->sync->suppressed_at_accumulation, (std::vector{false, false})); + } +} +TEST(PipelineScheduleTest, ExplicitStepRejectsMalformedBatchesBeforeOptimizerOrCommunication) { + auto chunk = std::make_shared(2.0f); + auto stage = std::make_shared(0, 1, std::vector>{{1}}, Device(), + std::vector>{chunk}); + const std::vector specs{{0, 1}}; + auto layout = std::make_shared(PipelineLayout::BuildChunkLayout(1, 1, specs)); + PipelineSchedule schedule(stage, 1, 2, layout, true); + const float values[]{1, 2, 3}; + auto valid = std::make_shared(values, std::vector{2}, DataType::kFLOAT32, Device()); + auto uneven = std::make_shared(values, std::vector{3}, DataType::kFLOAT32, Device()); + auto scalar = std::make_shared(values, std::vector{}, DataType::kFLOAT32, Device()); + // Null optimizer is deliberate: each validation must happen before ZeroGrad. + for (const auto &invalid : {std::shared_ptr{}, scalar, uneven}) { + EXPECT_THROW(schedule.Step(invalid, valid, nullptr, nullptr, DataType::kFLOAT32), std::invalid_argument); + EXPECT_THROW(schedule.Step(valid, invalid, nullptr, nullptr, DataType::kFLOAT32), std::invalid_argument); + } + EXPECT_EQ(chunk->sync->entries, 0); + EXPECT_EQ(chunk->parameter("weight")->grad(), nullptr); +} + +} // namespace +} // namespace infini_train::nn::parallel diff --git a/tests/distributed/test_pp_layout_parser.cc b/tests/distributed/test_pp_layout_parser.cc new file mode 100644 index 000000000..27eea53bc --- /dev/null +++ b/tests/distributed/test_pp_layout_parser.cc @@ -0,0 +1,115 @@ +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include + +#include "example/common/parser.h" + +namespace infini_train::examples { +namespace { +using nn::parallel::LayerIndex; +using nn::parallel::PipelineChunkSpec; + +TEST(PipelineLayoutCLITest, ResolvesDefaultAndInterleavedRequests) { + auto request = ParsePipelineLayoutRequest("", "", 2, 2); + EXPECT_TRUE(std::holds_alternative(request)); + EXPECT_FALSE(IsExplicitChunkRequest(request)); + const auto uniform = ResolvePipelineLayout(13, 2, 2, request); + ASSERT_NE(uniform, nullptr); + EXPECT_FALSE(uniform->IsCustom()); + EXPECT_EQ(uniform->GetNumChunks(), 4); + EXPECT_EQ(uniform->GetChunk(0).layer_range.end, 4); + + request = ParsePipelineLayoutRequest("02,4,3,3", "", 2, 2); + EXPECT_EQ(std::get>(request), (std::vector{2, 4, 3, 3})); + EXPECT_FALSE(IsExplicitChunkRequest(request)); + const auto custom = ResolvePipelineLayout(12, 2, 2, request); + ASSERT_NE(custom, nullptr); + EXPECT_TRUE(custom->IsCustom()); + EXPECT_EQ(custom->GetChunk(2).stage_id, 0); + EXPECT_EQ(custom->GetChunk(2).local_chunk_idx, 1); + EXPECT_EQ(custom->GetChunk(2).layer_range.begin, 6); + EXPECT_EQ(custom->GetChunk(2).layer_range.end, 9); + EXPECT_THROW((void)ResolvePipelineLayout(13, 2, 2, request), std::invalid_argument); +} + +TEST(PipelineLayoutCLITest, PreservesExplicitOwnerOrderAndValidatesActualModelSize) { + const auto request = ParsePipelineLayoutRequest("", "1:2,0:3,1:7", 2, 2); + ASSERT_TRUE(IsExplicitChunkRequest(request)); + const auto &specs = std::get>(request); + ASSERT_EQ(specs.size(), 3); + EXPECT_EQ(specs[0].stage_id, 1); + EXPECT_EQ(specs[0].layer_count, 2); + EXPECT_EQ(specs[1].stage_id, 0); + EXPECT_EQ(specs[2].layer_count, 7); + const auto layout = ResolvePipelineLayout(12, 2, 2, request); + ASSERT_NE(layout, nullptr); + EXPECT_EQ(layout->GetNumChunks(), 3); // Not PP * vPP. + EXPECT_EQ(layout->GetMaxLocalChunks(), 2); + EXPECT_EQ(layout->GetInputStage(), 1); + EXPECT_EQ(layout->GetOutputStage(), 1); + EXPECT_EQ(layout->GetChunk(2).layer_range.begin, 5); + EXPECT_THROW((void)ResolvePipelineLayout(13, 2, 2, request), std::invalid_argument); + EXPECT_THROW((void)ResolvePipelineLayout(12, 2, 1, request), std::invalid_argument); +} + +TEST(PipelineLayoutCLITest, FallsBackOnlyForUnderfilledDefaults) { + const auto request = ParsePipelineLayoutRequest("", "", 2, 2); + EXPECT_EQ(ResolvePipelineLayout(0, 2, 2, request), nullptr); + EXPECT_EQ(ResolvePipelineLayout(3, 2, 2, request), nullptr); + EXPECT_NE(ResolvePipelineLayout(4, 2, 2, request), nullptr); + EXPECT_THROW((void)ResolvePipelineLayout(-1, 2, 2, request), std::invalid_argument); + EXPECT_THROW((void)ResolvePipelineLayout(0, 2, 2, PipelineLayoutRequest{std::vector{}}), + std::invalid_argument); + EXPECT_THROW((void)ResolvePipelineLayout(0, 2, 2, PipelineLayoutRequest{std::vector{}}), + std::invalid_argument); + const PipelineLayoutRequest explicit_request = std::vector{{0, 2}, {1, 2}}; + EXPECT_THROW((void)ResolvePipelineLayout(4, 0, 1, explicit_request), std::invalid_argument); + EXPECT_THROW((void)ResolvePipelineLayout(4, 2, 0, explicit_request), std::invalid_argument); + EXPECT_THROW((void)ResolvePipelineLayout(4, 2, std::numeric_limits::max(), explicit_request), + std::invalid_argument); +} + +TEST(PipelineLayoutCLITest, RejectsMalformedAndConflictingCLI) { + // One representative per syntax/semantic error class, rather than a Cartesian product. + for (const auto *value : {"1,", ",1", "1,,2", "1, 2", "+1,2", "-1,2", "0,2", "1.5,2", "9223372036854775808,1"}) { + SCOPED_TRACE(value); + EXPECT_THROW((void)ParsePipelineLayerPartition(value), std::invalid_argument); + } + for (const auto *value : {"0:6,1:6,", "0:6:1,1:6", "0:,1:12", "0:6, 1:6", "0:0,1:12", "0:6,2:6", "0:6,0:6", + "-1:6,1:6", "2147483648:6,1:6", "0:9223372036854775808,1:1"}) { + SCOPED_TRACE(value); + EXPECT_THROW((void)ParsePipelineLayoutRequest("", value, 2, 1), std::invalid_argument); + } + EXPECT_THROW((void)ParsePipelineLayoutRequest("6,6", "0:6,1:6", 2, 1), std::invalid_argument); + EXPECT_THROW((void)ParsePipelineLayoutRequest("6,6", "", 2, 2), std::invalid_argument); + EXPECT_THROW((void)ParsePipelineLayoutRequest("", "0:3,1:3,0:6", 2, 1), std::invalid_argument); + EXPECT_THROW((void)ParsePipelineLayoutRequest("", "", 0, 1), std::invalid_argument); + EXPECT_THROW((void)ParsePipelineLayoutRequest("", "", 2, 0), std::invalid_argument); + EXPECT_THROW((void)ParsePipelineLayoutRequest("", "", 2, std::numeric_limits::max()), std::invalid_argument); + try { + (void)ParsePipelineLayerPartition("6,x"); + FAIL() << "Malformed count must be rejected"; + } catch (const std::invalid_argument &error) { + const std::string message = error.what(); + EXPECT_NE(message.find("pipeline_layer_partition"), std::string::npos); + EXPECT_NE(message.find("6,x"), std::string::npos); + EXPECT_NE(message.find("item 1"), std::string::npos); + } +} + +TEST(PipelineLayoutCLITest, FormatsActualChunkIdentityAndEndpoints) { + const auto layout = ResolvePipelineLayout(12, 2, 2, ParsePipelineLayoutRequest("", "1:2,0:3,1:7", 2, 2)); + const auto formatted = FormatPipelineLayout(*layout); + EXPECT_NE(formatted.find("custom pipeline layout:"), std::string::npos); + EXPECT_NE(formatted.find("Stage 0: chunk 1 (local 0) layers [2,5)"), std::string::npos); + EXPECT_NE( + formatted.find( + "Stage 1: chunk 0 (local 0) layers [0,2) chunk 2 (local 1) layers [5,12) embedding final_norm lm_head"), + std::string::npos); +} +} // namespace +} // namespace infini_train::examples