From a00e8b954110f74e16700e0a0ea40a2d3c986b50 Mon Sep 17 00:00:00 2001 From: cx Date: Fri, 3 Jul 2026 09:11:44 +0000 Subject: [PATCH 1/8] Support torchrun-style InfiniTrain multi-process launch --- example/gpt2/main.cc | 7 +- example/llama3/main.cc | 7 +- infini_train/include/nn/parallel/global.h | 1 + infini_train/src/device.cc | 11 +- infini_train/src/nn/parallel/data_parallel.cc | 6 +- .../parallel/ddp/distributed_data_parallel.cc | 6 +- infini_train/src/nn/parallel/global.cc | 29 ++++- infini_train/src/nn/parallel/process_group.cc | 2 +- scripts/run_models_and_profile.bash | 20 ++- scripts/test_config.json | 120 ++++++++++++------ tools/infini_run/infini_run.cc | 100 +++++++++++++-- 11 files changed, 238 insertions(+), 71 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index fd2522e3..7563c7e2 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -180,7 +180,7 @@ void Train(const nn::parallel::Rank &rank) { const ProcessGroup *pp_pg = nullptr; if (rank.IsParallel()) { - device = Device(Device::DeviceType::kCUDA, rank.thread_rank()); + device = Device(Device::DeviceType::kCUDA, global::GetLocalDeviceIndex(rank.thread_rank())); auto *pg_factory = ProcessGroupFactory::Instance(device.type()); if (ddp_world_size > 1) { @@ -565,7 +565,6 @@ int main(int argc, char *argv[]) { LOG(INFO) << nn::parallel::global::ProcessGroupOverview(); - // NOTE(dcj): currently we only support single process if (FLAGS_nthread_per_process > 1) { std::vector threads; for (int idx = 0; idx < FLAGS_nthread_per_process; ++idx) { @@ -576,7 +575,9 @@ int main(int argc, char *argv[]) { for (auto &thread : threads) { thread.join(); } } else { - Train({0, 0, 1, 1}); + nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), + FLAGS_nthread_per_process); + Train(rank); } gflags::ShutDownCommandLineFlags(); diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 05fad4a8..2510a8f5 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -168,7 +168,7 @@ void Train(const nn::parallel::Rank &rank) { const ProcessGroup *pp_pg = nullptr; if (rank.IsParallel()) { - device = Device(Device::DeviceType::kCUDA, rank.thread_rank()); + device = Device(Device::DeviceType::kCUDA, global::GetLocalDeviceIndex(rank.thread_rank())); auto *pg_factory = ProcessGroupFactory::Instance(device.type()); if (ddp_world_size > 1) { @@ -544,7 +544,6 @@ int main(int argc, char *argv[]) { LOG(INFO) << nn::parallel::global::ProcessGroupOverview(); - // NOTE(dcj): currently we only support single process if (FLAGS_nthread_per_process > 1) { std::vector threads; for (int idx = 0; idx < FLAGS_nthread_per_process; ++idx) { @@ -555,7 +554,9 @@ int main(int argc, char *argv[]) { for (auto &thread : threads) { thread.join(); } } else { - Train({0, 0, 1, 1}); + nn::parallel::Rank rank(nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), + FLAGS_nthread_per_process); + Train(rank); } gflags::ShutDownCommandLineFlags(); diff --git a/infini_train/include/nn/parallel/global.h b/infini_train/include/nn/parallel/global.h index 9373100f..c7c74a99 100644 --- a/infini_train/include/nn/parallel/global.h +++ b/infini_train/include/nn/parallel/global.h @@ -98,6 +98,7 @@ inline int GetNprocPerNode() { return GlobalEnv::Instance().nproc_per_node(); } inline int GetNthreadPerProc() { return GlobalEnv::Instance().nthread_per_process(); } inline int GetGlobalProcRank() { return GlobalEnv::Instance().global_proc_rank(); } inline int GetLocalProcRank() { return GlobalEnv::Instance().local_proc_rank(); } +inline int GetLocalDeviceIndex(int thread_rank = 0) { return GetLocalProcRank() * GetNthreadPerProc() + thread_rank; } inline int GetTensorParallelSize() { return GlobalEnv::Instance().tensor_parallel_size(); } inline int GetSequenceParallelSize() { return GlobalEnv::Instance().sequence_parallel_size(); } diff --git a/infini_train/src/device.cc b/infini_train/src/device.cc index 1bb3aaad..db10cd54 100644 --- a/infini_train/src/device.cc +++ b/infini_train/src/device.cc @@ -33,7 +33,16 @@ std::string Device::ToString() const { } nn::parallel::Rank Device::Rank() const { - return {nn::parallel::global::GetGlobalProcRank(), index_, nn::parallel::global::GetNprocPerNode(), + if (IsCPU()) { + return {nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), + nn::parallel::global::GetNthreadPerProc()}; + } + + const int thread_rank = index_ - nn::parallel::global::GetLocalDeviceIndex(); + CHECK_GE(thread_rank, 0) << "CUDA device index is outside the current process rank range"; + CHECK_LT(thread_rank, nn::parallel::global::GetNthreadPerProc()) + << "CUDA device index is outside the current process rank range"; + return {nn::parallel::global::GetGlobalProcRank(), thread_rank, nn::parallel::global::GetNprocPerNode(), nn::parallel::global::GetNthreadPerProc()}; } diff --git a/infini_train/src/nn/parallel/data_parallel.cc b/infini_train/src/nn/parallel/data_parallel.cc index c48761b7..6199d39f 100644 --- a/infini_train/src/nn/parallel/data_parallel.cc +++ b/infini_train/src/nn/parallel/data_parallel.cc @@ -60,7 +60,11 @@ ParallelApply(const std::vector> &modules, DataParallel::DataParallel(const std::shared_ptr &module, int dim, Device::DeviceType device_type) : dim_(dim) { devices_.reserve(global::GetNthreadPerProc()); - for (int index = 0; index < global::GetNthreadPerProc(); ++index) { devices_.emplace_back(device_type, index); } + for (int thread_rank = 0; thread_rank < global::GetNthreadPerProc(); ++thread_rank) { + const int device_index + = device_type == Device::DeviceType::kCUDA ? global::GetLocalDeviceIndex(thread_rank) : thread_rank; + devices_.emplace_back(device_type, device_index); + } CHECK_GT(devices_.size(), 0) << "No available devices found"; output_device_ = devices_.at(0); diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index b08b64fb..0bdf4046 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -10,6 +10,7 @@ #include "infini_train/include/autograd/function_hook.h" #include "infini_train/include/nn/modules/module.h" +#include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/nn/parallel/parallel_functional.h" #include "infini_train/include/nn/parallel/process_group.h" #include "infini_train/include/nn/parallel/rank.h" @@ -32,7 +33,8 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod continue; } auto device = param->GetDevice(); - CHECK_EQ(device.index(), rank.thread_rank()) << "All parameters must be on the same device as the module"; + CHECK_EQ(device.index(), global::GetLocalDeviceIndex(rank.thread_rank())) + << "All parameters must be on the same device as the module"; if (!ddp_config.gradient_bucketing_enabled && ddp_config.zero_stage < 1) { auto hook = std::make_unique( function::ReduceOpType::kAvg, ddp_pg_); @@ -40,7 +42,7 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod } } for (auto &buffer : module->Buffers()) { - CHECK_EQ(buffer->GetDevice().index(), rank.thread_rank()) + CHECK_EQ(buffer->GetDevice().index(), global::GetLocalDeviceIndex(rank.thread_rank())) << "All buffers must be on the same device as the module"; } modules_[kModuleName] = std::move(module); diff --git a/infini_train/src/nn/parallel/global.cc b/infini_train/src/nn/parallel/global.cc index 65a3208e..e4399128 100644 --- a/infini_train/src/nn/parallel/global.cc +++ b/infini_train/src/nn/parallel/global.cc @@ -13,6 +13,8 @@ int GetEnvAsInt(const std::string &name, int default_value) { return value ? std::atoi(value) : default_value; } +bool HasEnv(const std::string &name) { return std::getenv(name.c_str()) != nullptr; } + } // namespace namespace infini_train::nn::parallel::global { @@ -92,13 +94,30 @@ void GlobalEnv::Init(int nthread_per_process, int tensor_parallel_size, bool seq CHECK(!initialized_) << "Repeated initialization of GlobalEnv!"; - nnodes_ = GetEnvAsInt("NNODES", 1); - nproc_per_node_ = GetEnvAsInt("NPROC_PER_NODE", 1); - world_size_ = GetEnvAsInt("PROC_WORLD_SIZE", 1) * nthread_per_process; - global_proc_rank_ = GetEnvAsInt("GLOBAL_PROC_RANK", 0); - local_proc_rank_ = GetEnvAsInt("LOCAL_PROC_RANK", 0); + const int proc_world_size = GetEnvAsInt("PROC_WORLD_SIZE", GetEnvAsInt("WORLD_SIZE", 1)); + nproc_per_node_ = GetEnvAsInt("NPROC_PER_NODE", GetEnvAsInt("LOCAL_WORLD_SIZE", 1)); + CHECK_GT(nproc_per_node_, 0) << "NPROC_PER_NODE/LOCAL_WORLD_SIZE must be positive"; + CHECK_GT(proc_world_size, 0) << "PROC_WORLD_SIZE/WORLD_SIZE must be positive"; + CHECK_EQ(proc_world_size % nproc_per_node_, 0) + << "PROC_WORLD_SIZE/WORLD_SIZE must be divisible by NPROC_PER_NODE/LOCAL_WORLD_SIZE"; + const bool nnodes_env_set = HasEnv("NNODES"); + nnodes_ = GetEnvAsInt("NNODES", proc_world_size / nproc_per_node_); + CHECK_GT(nnodes_, 0) << "NNODES must be positive"; + if (nnodes_env_set) { + CHECK_EQ(nnodes_ * nproc_per_node_, proc_world_size) + << "NNODES * NPROC_PER_NODE/LOCAL_WORLD_SIZE must equal PROC_WORLD_SIZE/WORLD_SIZE"; + } + global_proc_rank_ = GetEnvAsInt("GLOBAL_PROC_RANK", GetEnvAsInt("RANK", 0)); + local_proc_rank_ = GetEnvAsInt("LOCAL_PROC_RANK", GetEnvAsInt("LOCAL_RANK", 0)); + CHECK_GE(global_proc_rank_, 0) << "GLOBAL_PROC_RANK/RANK must be non-negative"; + CHECK_LT(global_proc_rank_, proc_world_size) + << "GLOBAL_PROC_RANK/RANK must be less than PROC_WORLD_SIZE/WORLD_SIZE"; + CHECK_GE(local_proc_rank_, 0) << "LOCAL_PROC_RANK/LOCAL_RANK must be non-negative"; + CHECK_LT(local_proc_rank_, nproc_per_node_) + << "LOCAL_PROC_RANK/LOCAL_RANK must be less than NPROC_PER_NODE/LOCAL_WORLD_SIZE"; nthread_per_process_ = nthread_per_process; + world_size_ = proc_world_size * nthread_per_process; CHECK_GE(tensor_parallel_size, 1) << "Tensor Parallel size must be >= 1"; tensor_parallel_size_ = tensor_parallel_size; sequence_parallel_enabled_ = sequence_parallel_enabled; diff --git a/infini_train/src/nn/parallel/process_group.cc b/infini_train/src/nn/parallel/process_group.cc index 174aa645..fda07e13 100644 --- a/infini_train/src/nn/parallel/process_group.cc +++ b/infini_train/src/nn/parallel/process_group.cc @@ -99,7 +99,7 @@ void ProcessGroup::InitMultiProcess(const std::vector &ranks) { int global_thread_rank = lower_rank + i; auto it = std::ranges::find(ranks, global_thread_rank); if (it != ranks.end()) { - auto device = Device(backend_, i); + auto device = Device(backend_, global::GetLocalDeviceIndex(i)); core::DeviceGuard guard(device); core::CclComm *comm_raw = nullptr; diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index b7a7b5cc..f3174ab1 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -439,6 +439,18 @@ check_model_inputs() { fi } +model_cmd_for_test() { + local model_bin="$1" + local input_bin="$2" + local llmc_filepath="$3" + local arg_str="$4" + local nproc_per_node="$5" + + printf './infini_run --nproc_per_node=%s %s --input_bin %q --llmc_filepath %q --device cuda %s' \ + "$nproc_per_node" "$model_bin" \ + "$input_bin" "$llmc_filepath" "$arg_str" +} + # Run tests num_basic_compile_commands=$(jq '.basic_compile_commands | length' "$CONFIG_FILE") num_groups=$(jq '.test_groups | length' "$CONFIG_FILE") @@ -510,23 +522,25 @@ for ((id=0; id BuildLauncherArgv(int train_program_index, char **argv) { + std::vector launcher_argv; + launcher_argv.push_back(argv[0]); + for (int i = 1; i < train_program_index; ++i) { launcher_argv.push_back(argv[i]); } + launcher_argv.push_back(nullptr); + return launcher_argv; +} + +void SetEnvInt(const char *name, int value) { + const auto value_str = std::to_string(value); + setenv(name, value_str.c_str(), 1); +} + +} // namespace + int main(int argc, char **argv) { - gflags::ParseCommandLineFlags(&argc, &argv, true); + const int train_program_index = FindTrainProgramIndex(argc, argv); + std::vector launcher_argv = BuildLauncherArgv(train_program_index, argv); + int launcher_argc = static_cast(launcher_argv.size()) - 1; + char **launcher_argv_ptr = launcher_argv.data(); + gflags::ParseCommandLineFlags(&launcher_argc, &launcher_argv_ptr, true); google::InitGoogleLogging(argv[0]); - CHECK_GE(argc, 2) << "No training prgram specified!"; + CHECK_GT(FLAGS_nnodes, 0) << "nnodes must be positive"; + CHECK_GT(FLAGS_nproc_per_node, 0) << "nproc_per_node must be positive"; + CHECK_NE(FLAGS_rdzv_endpoint.find(':'), std::string::npos) << "rdzv_endpoint must be host:port"; - std::string train_program = argv[1]; + CHECK_LT(train_program_index, argc) << "No training program specified!"; + + std::string train_program = argv[train_program_index]; + CHECK_NE(train_program, "--") << "Explicit '--' separator is not supported; pass the training program directly " + "after infini_run launcher flags"; std::vector train_argv; - for (int i = 1; i < argc; ++i) { train_argv.push_back(argv[i]); } + for (int i = train_program_index; i < argc; ++i) { train_argv.push_back(argv[i]); } train_argv.push_back(nullptr); - int world_size = FLAGS_nnodes * FLAGS_nproc_per_node; + int proc_world_size = FLAGS_nnodes * FLAGS_nproc_per_node; std::string master_addr = FLAGS_rdzv_endpoint.substr(0, FLAGS_rdzv_endpoint.find(':')); std::string master_port = FLAGS_rdzv_endpoint.substr(FLAGS_rdzv_endpoint.find(':') + 1); @@ -32,16 +82,23 @@ int main(int argc, char **argv) { pid_t pid = fork(); if (pid == 0) { int global_proc_rank = FLAGS_node_rank * FLAGS_nproc_per_node + local_proc_rank; - setenv("NNODES", std::to_string(FLAGS_nnodes).c_str(), 1); - setenv("NPROC_PER_NODE", std::to_string(FLAGS_nproc_per_node).c_str(), 1); + SetEnvInt("NNODES", FLAGS_nnodes); + SetEnvInt("NPROC_PER_NODE", FLAGS_nproc_per_node); + SetEnvInt("LOCAL_WORLD_SIZE", FLAGS_nproc_per_node); setenv("MASTER_ADDR", master_addr.c_str(), 1); setenv("MASTER_PORT", master_port.c_str(), 1); - setenv("GLOBAL_PROC_RANK", std::to_string(global_proc_rank).c_str(), 1); - setenv("LOCAL_PROC_RANK", std::to_string(local_proc_rank).c_str(), 1); + SetEnvInt("GLOBAL_PROC_RANK", global_proc_rank); + SetEnvInt("LOCAL_PROC_RANK", local_proc_rank); + SetEnvInt("RANK", global_proc_rank); + SetEnvInt("LOCAL_RANK", local_proc_rank); - setenv("PROC_WORLD_SIZE", std::to_string(world_size).c_str(), 1); + SetEnvInt("PROC_WORLD_SIZE", proc_world_size); + SetEnvInt("WORLD_SIZE", proc_world_size); + SetEnvInt("GROUP_RANK", FLAGS_node_rank); + SetEnvInt("ROLE_RANK", global_proc_rank); + SetEnvInt("ROLE_WORLD_SIZE", proc_world_size); execvp(train_program.c_str(), train_argv.data()); perror("exec failed"); @@ -49,10 +106,29 @@ int main(int argc, char **argv) { } } + int exit_code = 0; for (int i = 0; i < FLAGS_nproc_per_node; ++i) { int status; - wait(&status); + pid_t child = wait(&status); + if (child < 0) { + perror("wait failed"); + return 1; + } + + if (WIFEXITED(status)) { + int child_exit_code = WEXITSTATUS(status); + if (child_exit_code != 0 && exit_code == 0) { + exit_code = child_exit_code; + } + } else if (WIFSIGNALED(status)) { + int signal = WTERMSIG(status); + if (exit_code == 0) { + exit_code = 128 + signal; + } + } else if (exit_code == 0) { + exit_code = 1; + } } - return 0; + return exit_code; } From 5cbb4aae031b971d22ec258252f9f83b32c0cb21 Mon Sep 17 00:00:00 2001 From: cx Date: Thu, 9 Jul 2026 02:32:10 +0000 Subject: [PATCH 2/8] test: add basic 8-process config group Add a dedicated 8_proc test group containing the 8-process variants of the original basic multi-GPU cases. --- scripts/test_config.json | 255 ++++++++++++++++++++++++++------------- 1 file changed, 173 insertions(+), 82 deletions(-) diff --git a/scripts/test_config.json b/scripts/test_config.json index 6868bbd1..7cb778a2 100644 --- a/scripts/test_config.json +++ b/scripts/test_config.json @@ -15,8 +15,8 @@ "RUN_PROFILE_TEST": "true", "MIXTRAL_INPUT_BIN": "/data1/shared/InfiniTrain-dev/data/llmc/llama3/tinyshakespeare/tiny_shakespeare_train.bin", "MIXTRAL_LLMC_FILEPATH": "/data1/shared/InfiniTrain-dev/data/llmc/mixtral/mixtral_megatron_export.bin", - "GPT2_TEST_GROUPS": "basic,zero,lora,checkpoint,lr_scheduler", - "LLAMA3_TEST_GROUPS": "basic,zero,lora,checkpoint,lr_scheduler", + "GPT2_TEST_GROUPS": "basic,zero,lora,checkpoint,lr_scheduler,8_proc", + "LLAMA3_TEST_GROUPS": "basic,zero,lora,checkpoint,lr_scheduler,8_proc", "MIXTRAL_TEST_GROUPS": "moe" }, "basic_compile_commands": [ @@ -63,8 +63,7 @@ "id": "3", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120 @@ -74,8 +73,7 @@ "id": "3_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120 @@ -85,8 +83,7 @@ "id": "4", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -97,8 +94,7 @@ "id": "4_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -109,8 +105,7 @@ "id": "5", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -122,8 +117,7 @@ "id": "5_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -135,8 +129,7 @@ "id": "6", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -147,8 +140,7 @@ "id": "6_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -159,8 +151,7 @@ "id": "7", "args": { "dtype": "float32", - "nproc_per_node": 4, - "nthread_per_process": 1, + "nthread_per_process": 4, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -172,8 +163,7 @@ "id": "7_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 4, - "nthread_per_process": 1, + "nthread_per_process": 4, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -185,8 +175,7 @@ "id": "8", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -200,8 +189,7 @@ "id": "8_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -220,8 +208,7 @@ "id": "3_distopt", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -232,8 +219,7 @@ "id": "3_zero2", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -244,8 +230,7 @@ "id": "3_bfloat16_distopt", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -256,8 +241,7 @@ "id": "3_bfloat16_zero2", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -268,8 +252,7 @@ "id": "4_distopt", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -281,8 +264,7 @@ "id": "4_zero2", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -294,8 +276,7 @@ "id": "4_bfloat16_distopt", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -307,8 +288,7 @@ "id": "4_bfloat16_zero2", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -320,8 +300,7 @@ "id": "5_distopt", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -334,8 +313,7 @@ "id": "5_zero2", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -348,8 +326,7 @@ "id": "5_bfloat16_distopt", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -362,8 +339,7 @@ "id": "5_bfloat16_zero2", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -376,8 +352,7 @@ "id": "8_distopt", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -392,8 +367,7 @@ "id": "8_bfloat16_distopt", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -635,8 +609,7 @@ "id": "3_lora", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -650,8 +623,7 @@ "id": "3_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -665,8 +637,7 @@ "id": "4_lora", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -681,8 +652,7 @@ "id": "4_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -697,8 +667,7 @@ "id": "5_lora", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -713,8 +682,7 @@ "id": "5_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -729,8 +697,7 @@ "id": "6_lora", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -744,8 +711,7 @@ "id": "6_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -759,8 +725,7 @@ "id": "7_lora", "args": { "dtype": "float32", - "nproc_per_node": 4, - "nthread_per_process": 1, + "nthread_per_process": 4, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -775,8 +740,7 @@ "id": "7_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 4, - "nthread_per_process": 1, + "nthread_per_process": 4, "num_iteration": 10, "batch_size": 10, "total_batch_size": 5120, @@ -791,8 +755,7 @@ "id": "8_lora", "args": { "dtype": "float32", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -809,8 +772,7 @@ "id": "8_lora_bfloat16", "args": { "dtype": "bfloat16", - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 10, "batch_size": 40, "total_batch_size": 5120, @@ -867,8 +829,7 @@ { "id": "ckpt_3d_ddp_tp_pp", "args": { - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 50, "save_interval": 10, "batch_size": 40, @@ -884,8 +845,7 @@ { "id": "ckpt_3d_ddp_tp_pp_resume", "args": { - "nproc_per_node": 8, - "nthread_per_process": 1, + "nthread_per_process": 8, "num_iteration": 50, "save_interval": 10, "batch_size": 40, @@ -935,6 +895,137 @@ } } ] + }, + { + "tag": "8_proc", + "tests": [ + { + "id": "3", + "args": { + "dtype": "float32", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 10, + "total_batch_size": 5120 + } + }, + { + "id": "3_bfloat16", + "args": { + "dtype": "bfloat16", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 10, + "total_batch_size": 5120 + } + }, + { + "id": "4", + "args": { + "dtype": "float32", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 4 + } + }, + { + "id": "4_bfloat16", + "args": { + "dtype": "bfloat16", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 4 + } + }, + { + "id": "5", + "args": { + "dtype": "float32", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 4, + "sequence_parallel": true + } + }, + { + "id": "5_bfloat16", + "args": { + "dtype": "bfloat16", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 4, + "sequence_parallel": true + } + }, + { + "id": "6", + "args": { + "dtype": "float32", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 10, + "total_batch_size": 5120, + "pipeline_parallel": 8 + } + }, + { + "id": "6_bfloat16", + "args": { + "dtype": "bfloat16", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 10, + "total_batch_size": 5120, + "pipeline_parallel": 8 + } + }, + { + "id": "8", + "args": { + "dtype": "float32", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 2, + "sequence_parallel": true, + "pipeline_parallel": 2, + "virtual_pipeline_parallel": 2 + } + }, + { + "id": "8_bfloat16", + "args": { + "dtype": "bfloat16", + "nproc_per_node": 8, + "nthread_per_process": 1, + "num_iteration": 10, + "batch_size": 40, + "total_batch_size": 5120, + "tensor_parallel": 2, + "sequence_parallel": true, + "pipeline_parallel": 2, + "virtual_pipeline_parallel": 2 + } + } + ] } ] } From aa110a3ef6e473e06910455c00a8fd910a35d174 Mon Sep 17 00:00:00 2001 From: cx Date: Fri, 10 Jul 2026 11:08:50 +0000 Subject: [PATCH 3/8] fix: stabilize distributed training resume flow Track DataLoader progress by global batches so distributed ranks slice data consistently and can resume/cycle from saved consumption counts. Also scope CCL unique ID files per run, generate NCCL IDs only on the main rank, clean up run-local rendezvous files, and add DataLoader coverage. --- README.md | 5 +-- infini_train/include/core/ccl/ccl.h | 2 +- infini_train/src/core/ccl/ccl.cc | 4 ++- infini_train/src/core/ccl/ccl_utils.cc | 28 +++++++++++---- infini_train/src/core/ccl/cuda/nccl_impl.cc | 6 ++-- infini_train/src/core/ccl/cuda/nccl_impl.h | 2 +- infini_train/src/nn/parallel/process_group.cc | 8 +++-- tools/infini_run/infini_run.cc | 35 ++++++++++++++++++- 8 files changed, 73 insertions(+), 17 deletions(-) diff --git a/README.md b/README.md index 1a51d1a8..befab7a0 100644 --- a/README.md +++ b/README.md @@ -121,8 +121,9 @@ The following examples demonstrate **LLaMA 3 supervised fine-tuning (SFT)** usin ./infini_run \ --nnodes=2 \ --nproc_per_node=1 \ - --node_rank=[rank_id] \ - -- ./llama3 \ + --node_rank=[rank_id] \ + --rdzv_endpoint=[master_addr]:29500 \ + ./llama3 \ --device cuda \ --input_bin [training_data_path] \ --llmc_filepath [model_path] \ diff --git a/infini_train/include/core/ccl/ccl.h b/infini_train/include/core/ccl/ccl.h index 626cb078..5cce524d 100644 --- a/infini_train/include/core/ccl/ccl.h +++ b/infini_train/include/core/ccl/ccl.h @@ -28,7 +28,7 @@ class CclImpl { virtual void GetAsyncError(const CclComm *comm, CclStatus *async_error) const; - virtual void GetUniqueId(CclUniqueId **unique_id) const; + virtual void CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const; virtual void CommInitAll(CclComm **comms, int ndev, const int *devlist) const; diff --git a/infini_train/src/core/ccl/ccl.cc b/infini_train/src/core/ccl/ccl.cc index 92c14cc6..1c0f2b6e 100644 --- a/infini_train/src/core/ccl/ccl.cc +++ b/infini_train/src/core/ccl/ccl.cc @@ -16,7 +16,9 @@ void CclImpl::GetAsyncError(const CclComm *comm, CclStatus *async_error) const { LOG(FATAL) << "CclImpl::GetAsyncError is not implemented."; } -void CclImpl::GetUniqueId(CclUniqueId **unique_id) const { LOG(FATAL) << "CclImpl::GetUniqueId is not implemented."; } +void CclImpl::CreateUniqueId(CclUniqueId **, bool) const { + LOG(FATAL) << "CclImpl::CreateUniqueId is not implemented."; +} void CclImpl::CommInitAll(CclComm **comms, int ndev, const int *devlist) const { LOG(FATAL) << "CclImpl::CommInitAll is not implemented."; diff --git a/infini_train/src/core/ccl/ccl_utils.cc b/infini_train/src/core/ccl/ccl_utils.cc index f8968c6b..ebf1ed19 100644 --- a/infini_train/src/core/ccl/ccl_utils.cc +++ b/infini_train/src/core/ccl/ccl_utils.cc @@ -2,6 +2,7 @@ #include #include +#include #include #include #include @@ -11,13 +12,21 @@ namespace infini_train::core { namespace { -std::string UniqueIdFileName(const std::string &name, bool tmp = false) { - return "cclUniqueId_" + name + (tmp ? ".tmp" : ".bin"); +std::string UniqueIdPath(const std::string &pg_name) { + const char *run_id = std::getenv("INFINI_RUN_ID"); + const std::string prefix = run_id == nullptr ? "" : std::string(run_id) + "_"; + return "cclUniqueId_" + prefix + pg_name + ".bin"; +} + +std::string UniqueIdTmpPath(const std::string &pg_name) { + const char *run_id = std::getenv("INFINI_RUN_ID"); + const std::string prefix = run_id == nullptr ? "" : std::string(run_id) + "_"; + return "cclUniqueId_" + prefix + pg_name + ".tmp"; } } // namespace void WriteUniqueIdFile(const CclUniqueId &unique_id, const std::string &pg_name) { - const std::string tmp_path = UniqueIdFileName(pg_name, true); + const std::string tmp_path = UniqueIdTmpPath(pg_name); std::ofstream ofs(tmp_path, std::ios::binary); CHECK(ofs.good()) << "Failed to open unique_id tmp file for write: " << tmp_path; @@ -25,12 +34,14 @@ void WriteUniqueIdFile(const CclUniqueId &unique_id, const std::string &pg_name) ofs.write(reinterpret_cast(unique_id.Data()), static_cast(size)); ofs.close(); - std::rename(tmp_path.c_str(), UniqueIdFileName(pg_name).c_str()); + const std::string file_path = UniqueIdPath(pg_name); + CHECK_EQ(std::rename(tmp_path.c_str(), file_path.c_str()), 0) + << "Failed to rename unique_id file from " << tmp_path << " to " << file_path; } void ReadUniqueIdFile(CclUniqueId *unique_id, const std::string &pg_name) { CHECK_NOTNULL(unique_id); - const std::string file_path = UniqueIdFileName(pg_name); + const std::string file_path = UniqueIdPath(pg_name); while (!std::filesystem::exists(file_path)) { std::this_thread::sleep_for(std::chrono::microseconds(1000)); } @@ -46,10 +57,15 @@ void ReadUniqueIdFile(CclUniqueId *unique_id, const std::string &pg_name) { } void CleanupUniqueIdFile(const std::string &pg_name) { - const std::string file_path = UniqueIdFileName(pg_name); + const std::string file_path = UniqueIdPath(pg_name); if (std::filesystem::exists(file_path)) { std::filesystem::remove(file_path); } + + const std::string tmp_path = UniqueIdTmpPath(pg_name); + if (std::filesystem::exists(tmp_path)) { + std::filesystem::remove(tmp_path); + } } } // namespace infini_train::core diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.cc b/infini_train/src/core/ccl/cuda/nccl_impl.cc index 9e4b1a0d..d17cee46 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.cc +++ b/infini_train/src/core/ccl/cuda/nccl_impl.cc @@ -68,14 +68,16 @@ void NcclImpl::GetAsyncError(const CclComm *comm, CclStatus *async_error) const } } -void NcclImpl::GetUniqueId(CclUniqueId **unique_id) const { +void NcclImpl::CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const { CHECK_NOTNULL(unique_id); if (*unique_id == nullptr) { *unique_id = new NcclUniqueId(); } auto *nccl_unique_id = dynamic_cast(*unique_id); CHECK_NOTNULL(nccl_unique_id); - NCCL_CHECK(ncclGetUniqueId(nccl_unique_id->nccl_unique_id())); + if (generate_id) { + NCCL_CHECK(ncclGetUniqueId(nccl_unique_id->nccl_unique_id())); + } } void NcclImpl::CommInitAll(CclComm **comms, int ndev, const int *devlist) const { diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.h b/infini_train/src/core/ccl/cuda/nccl_impl.h index fca177fd..6fda67c4 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.h +++ b/infini_train/src/core/ccl/cuda/nccl_impl.h @@ -17,7 +17,7 @@ class NcclImpl final : public CclImpl { void GetAsyncError(const CclComm *comm, CclStatus *async_error) const override; - void GetUniqueId(CclUniqueId **unique_id) const override; + void CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const override; void CommInitAll(CclComm **comms, int ndev, const int *devlist) const override; diff --git a/infini_train/src/nn/parallel/process_group.cc b/infini_train/src/nn/parallel/process_group.cc index fda07e13..ad228df5 100644 --- a/infini_train/src/nn/parallel/process_group.cc +++ b/infini_train/src/nn/parallel/process_group.cc @@ -83,11 +83,13 @@ void ProcessGroup::InitMultiProcess(const std::vector &ranks) { int upper_rank = (global_proc_rank + 1) * n_threads; core::CclUniqueId *unique_id_raw = nullptr; - ccl_impl_->GetUniqueId(&unique_id_raw); - std::unique_ptr unique_id(unique_id_raw); int min_rank = std::ranges::min(ranks); - if (min_rank < upper_rank && min_rank >= lower_rank) { + bool is_main_rank = min_rank < upper_rank && min_rank >= lower_rank; + ccl_impl_->CreateUniqueId(&unique_id_raw, is_main_rank); + std::unique_ptr unique_id(unique_id_raw); + + if (is_main_rank) { is_main_process_ = true; core::WriteUniqueIdFile(*unique_id, name_); } else { diff --git a/tools/infini_run/infini_run.cc b/tools/infini_run/infini_run.cc index b632f64d..189d2752 100644 --- a/tools/infini_run/infini_run.cc +++ b/tools/infini_run/infini_run.cc @@ -1,7 +1,10 @@ +#include #include #include +#include #include #include +#include #include #include @@ -51,6 +54,28 @@ void SetEnvInt(const char *name, int value) { setenv(name, value_str.c_str(), 1); } +void CleanupRunUniqueIdFiles(const std::string &run_id) { + const std::string prefix = "cclUniqueId_" + run_id + "_"; + for (const auto &entry : std::filesystem::directory_iterator(std::filesystem::current_path())) { + if (!entry.is_regular_file()) { + continue; + } + const std::string filename = entry.path().filename().string(); + if (filename.rfind(prefix, 0) == 0) { + std::error_code ec; + std::filesystem::remove(entry.path(), ec); + if (ec) { + LOG(WARNING) << "Failed to remove unique-id file " << entry.path() << ": " << ec.message(); + } + } + } +} + +std::string GenerateLocalRunId() { + const auto now = std::chrono::steady_clock::now().time_since_epoch().count(); + return std::to_string(getpid()) + "_" + std::to_string(now); +} + } // namespace int main(int argc, char **argv) { @@ -77,6 +102,7 @@ int main(int argc, char **argv) { int proc_world_size = FLAGS_nnodes * FLAGS_nproc_per_node; std::string master_addr = FLAGS_rdzv_endpoint.substr(0, FLAGS_rdzv_endpoint.find(':')); std::string master_port = FLAGS_rdzv_endpoint.substr(FLAGS_rdzv_endpoint.find(':') + 1); + const std::string run_id = FLAGS_nnodes == 1 ? GenerateLocalRunId() : ""; for (int local_proc_rank = 0; local_proc_rank < FLAGS_nproc_per_node; ++local_proc_rank) { pid_t pid = fork(); @@ -88,6 +114,9 @@ int main(int argc, char **argv) { setenv("MASTER_ADDR", master_addr.c_str(), 1); setenv("MASTER_PORT", master_port.c_str(), 1); + if (!run_id.empty()) { + setenv("INFINI_RUN_ID", run_id.c_str(), 1); + } SetEnvInt("GLOBAL_PROC_RANK", global_proc_rank); SetEnvInt("LOCAL_PROC_RANK", local_proc_rank); @@ -112,7 +141,8 @@ int main(int argc, char **argv) { pid_t child = wait(&status); if (child < 0) { perror("wait failed"); - return 1; + exit_code = 1; + break; } if (WIFEXITED(status)) { @@ -130,5 +160,8 @@ int main(int argc, char **argv) { } } + if (!run_id.empty()) { + CleanupRunUniqueIdFiles(run_id); + } return exit_code; } From fb223c8b5149a33bb8526de45802615694b2482b Mon Sep 17 00:00:00 2001 From: cx Date: Thu, 30 Jul 2026 03:33:27 +0000 Subject: [PATCH 4/8] fix: correct distributed rank and unique ID handling - derive parallel state from the global world size - clarify global rank and per-node process semantics - add multi-node rank regression coverage - restore the NCCL-compatible GetUniqueId interface --- infini_train/include/core/ccl/ccl.h | 2 +- infini_train/include/nn/parallel/rank.h | 10 +-- infini_train/src/core/ccl/ccl.cc | 4 +- infini_train/src/core/ccl/cuda/nccl_impl.cc | 6 +- infini_train/src/core/ccl/cuda/nccl_impl.h | 2 +- infini_train/src/nn/parallel/global.cc | 18 ++---- infini_train/src/nn/parallel/process_group.cc | 8 +-- infini_train/src/nn/parallel/rank.cc | 15 ++--- tests/CMakeLists.txt | 3 + tests/distributed/CMakeLists.txt | 25 ++++++++ tests/distributed/test_rank.cc | 16 +++++ tools/infini_run/infini_run.cc | 63 ++++++++++++------- 12 files changed, 110 insertions(+), 62 deletions(-) create mode 100644 tests/distributed/CMakeLists.txt create mode 100644 tests/distributed/test_rank.cc diff --git a/infini_train/include/core/ccl/ccl.h b/infini_train/include/core/ccl/ccl.h index 5cce524d..626cb078 100644 --- a/infini_train/include/core/ccl/ccl.h +++ b/infini_train/include/core/ccl/ccl.h @@ -28,7 +28,7 @@ class CclImpl { virtual void GetAsyncError(const CclComm *comm, CclStatus *async_error) const; - virtual void CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const; + virtual void GetUniqueId(CclUniqueId **unique_id) const; virtual void CommInitAll(CclComm **comms, int ndev, const int *devlist) const; diff --git a/infini_train/include/nn/parallel/rank.h b/infini_train/include/nn/parallel/rank.h index eb61417d..baa623db 100644 --- a/infini_train/include/nn/parallel/rank.h +++ b/infini_train/include/nn/parallel/rank.h @@ -3,7 +3,7 @@ namespace infini_train::nn::parallel { class Rank { public: - Rank(int process_rank, int thread_rank, int process_size, int thread_size); + Rank(int global_process_rank, int thread_rank, int processes_per_node, int threads_per_process); int process_rank() const; int thread_rank() const; @@ -19,9 +19,9 @@ class Rank { bool IsLastRank() const; private: - const int process_rank_ = 0; // Rank of the current process within the node - const int thread_rank_ = 0; // Rank of the current thread within the process - const int process_size_ = 1; // Total number of processes on this node - const int thread_size_ = 1; // Total number of threads in the current process + const int global_process_rank_ = 0; // Global rank of the current process + const int thread_rank_ = 0; // Rank of the current thread within the process + const int processes_per_node_ = 1; // Number of processes on each node + const int threads_per_process_ = 1; // Number of threads in each process }; } // namespace infini_train::nn::parallel diff --git a/infini_train/src/core/ccl/ccl.cc b/infini_train/src/core/ccl/ccl.cc index 1c0f2b6e..92c14cc6 100644 --- a/infini_train/src/core/ccl/ccl.cc +++ b/infini_train/src/core/ccl/ccl.cc @@ -16,9 +16,7 @@ void CclImpl::GetAsyncError(const CclComm *comm, CclStatus *async_error) const { LOG(FATAL) << "CclImpl::GetAsyncError is not implemented."; } -void CclImpl::CreateUniqueId(CclUniqueId **, bool) const { - LOG(FATAL) << "CclImpl::CreateUniqueId is not implemented."; -} +void CclImpl::GetUniqueId(CclUniqueId **unique_id) const { LOG(FATAL) << "CclImpl::GetUniqueId is not implemented."; } void CclImpl::CommInitAll(CclComm **comms, int ndev, const int *devlist) const { LOG(FATAL) << "CclImpl::CommInitAll is not implemented."; diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.cc b/infini_train/src/core/ccl/cuda/nccl_impl.cc index d17cee46..9e4b1a0d 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.cc +++ b/infini_train/src/core/ccl/cuda/nccl_impl.cc @@ -68,16 +68,14 @@ void NcclImpl::GetAsyncError(const CclComm *comm, CclStatus *async_error) const } } -void NcclImpl::CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const { +void NcclImpl::GetUniqueId(CclUniqueId **unique_id) const { CHECK_NOTNULL(unique_id); if (*unique_id == nullptr) { *unique_id = new NcclUniqueId(); } auto *nccl_unique_id = dynamic_cast(*unique_id); CHECK_NOTNULL(nccl_unique_id); - if (generate_id) { - NCCL_CHECK(ncclGetUniqueId(nccl_unique_id->nccl_unique_id())); - } + NCCL_CHECK(ncclGetUniqueId(nccl_unique_id->nccl_unique_id())); } void NcclImpl::CommInitAll(CclComm **comms, int ndev, const int *devlist) const { diff --git a/infini_train/src/core/ccl/cuda/nccl_impl.h b/infini_train/src/core/ccl/cuda/nccl_impl.h index 6fda67c4..fca177fd 100644 --- a/infini_train/src/core/ccl/cuda/nccl_impl.h +++ b/infini_train/src/core/ccl/cuda/nccl_impl.h @@ -17,7 +17,7 @@ class NcclImpl final : public CclImpl { void GetAsyncError(const CclComm *comm, CclStatus *async_error) const override; - void CreateUniqueId(CclUniqueId **unique_id, bool generate_id) const override; + void GetUniqueId(CclUniqueId **unique_id) const override; void CommInitAll(CclComm **comms, int ndev, const int *devlist) const override; diff --git a/infini_train/src/nn/parallel/global.cc b/infini_train/src/nn/parallel/global.cc index e4399128..e56cac72 100644 --- a/infini_train/src/nn/parallel/global.cc +++ b/infini_train/src/nn/parallel/global.cc @@ -13,8 +13,6 @@ int GetEnvAsInt(const std::string &name, int default_value) { return value ? std::atoi(value) : default_value; } -bool HasEnv(const std::string &name) { return std::getenv(name.c_str()) != nullptr; } - } // namespace namespace infini_train::nn::parallel::global { @@ -94,21 +92,15 @@ void GlobalEnv::Init(int nthread_per_process, int tensor_parallel_size, bool seq CHECK(!initialized_) << "Repeated initialization of GlobalEnv!"; - const int proc_world_size = GetEnvAsInt("PROC_WORLD_SIZE", GetEnvAsInt("WORLD_SIZE", 1)); - nproc_per_node_ = GetEnvAsInt("NPROC_PER_NODE", GetEnvAsInt("LOCAL_WORLD_SIZE", 1)); + const int proc_world_size = GetEnvAsInt("WORLD_SIZE", GetEnvAsInt("PROC_WORLD_SIZE", 1)); + nproc_per_node_ = GetEnvAsInt("LOCAL_WORLD_SIZE", GetEnvAsInt("NPROC_PER_NODE", 1)); CHECK_GT(nproc_per_node_, 0) << "NPROC_PER_NODE/LOCAL_WORLD_SIZE must be positive"; CHECK_GT(proc_world_size, 0) << "PROC_WORLD_SIZE/WORLD_SIZE must be positive"; CHECK_EQ(proc_world_size % nproc_per_node_, 0) << "PROC_WORLD_SIZE/WORLD_SIZE must be divisible by NPROC_PER_NODE/LOCAL_WORLD_SIZE"; - const bool nnodes_env_set = HasEnv("NNODES"); - nnodes_ = GetEnvAsInt("NNODES", proc_world_size / nproc_per_node_); - CHECK_GT(nnodes_, 0) << "NNODES must be positive"; - if (nnodes_env_set) { - CHECK_EQ(nnodes_ * nproc_per_node_, proc_world_size) - << "NNODES * NPROC_PER_NODE/LOCAL_WORLD_SIZE must equal PROC_WORLD_SIZE/WORLD_SIZE"; - } - global_proc_rank_ = GetEnvAsInt("GLOBAL_PROC_RANK", GetEnvAsInt("RANK", 0)); - local_proc_rank_ = GetEnvAsInt("LOCAL_PROC_RANK", GetEnvAsInt("LOCAL_RANK", 0)); + nnodes_ = proc_world_size / nproc_per_node_; + global_proc_rank_ = GetEnvAsInt("RANK", GetEnvAsInt("GLOBAL_PROC_RANK", 0)); + local_proc_rank_ = GetEnvAsInt("LOCAL_RANK", GetEnvAsInt("LOCAL_PROC_RANK", 0)); CHECK_GE(global_proc_rank_, 0) << "GLOBAL_PROC_RANK/RANK must be non-negative"; CHECK_LT(global_proc_rank_, proc_world_size) << "GLOBAL_PROC_RANK/RANK must be less than PROC_WORLD_SIZE/WORLD_SIZE"; diff --git a/infini_train/src/nn/parallel/process_group.cc b/infini_train/src/nn/parallel/process_group.cc index ad228df5..fda07e13 100644 --- a/infini_train/src/nn/parallel/process_group.cc +++ b/infini_train/src/nn/parallel/process_group.cc @@ -83,13 +83,11 @@ void ProcessGroup::InitMultiProcess(const std::vector &ranks) { int upper_rank = (global_proc_rank + 1) * n_threads; core::CclUniqueId *unique_id_raw = nullptr; - - int min_rank = std::ranges::min(ranks); - bool is_main_rank = min_rank < upper_rank && min_rank >= lower_rank; - ccl_impl_->CreateUniqueId(&unique_id_raw, is_main_rank); + ccl_impl_->GetUniqueId(&unique_id_raw); std::unique_ptr unique_id(unique_id_raw); - if (is_main_rank) { + int min_rank = std::ranges::min(ranks); + if (min_rank < upper_rank && min_rank >= lower_rank) { is_main_process_ = true; core::WriteUniqueIdFile(*unique_id, name_); } else { diff --git a/infini_train/src/nn/parallel/rank.cc b/infini_train/src/nn/parallel/rank.cc index 617a8e50..d5b4a788 100644 --- a/infini_train/src/nn/parallel/rank.cc +++ b/infini_train/src/nn/parallel/rank.cc @@ -2,18 +2,19 @@ #include "infini_train/include/nn/parallel/global.h" namespace infini_train::nn::parallel { -Rank::Rank(int process_rank, int thread_rank, int process_size, int thread_size) - : process_rank_(process_rank), thread_rank_(thread_rank), process_size_(process_size), thread_size_(thread_size) {} +Rank::Rank(int global_process_rank, int thread_rank, int processes_per_node, int threads_per_process) + : global_process_rank_(global_process_rank), thread_rank_(thread_rank), processes_per_node_(processes_per_node), + threads_per_process_(threads_per_process) {} -int Rank::process_rank() const { return process_rank_; } +int Rank::process_rank() const { return global_process_rank_; } int Rank::thread_rank() const { return thread_rank_; } -int Rank::process_size() const { return process_size_; } -int Rank::thread_size() const { return thread_size_; } +int Rank::process_size() const { return processes_per_node_; } +int Rank::thread_size() const { return threads_per_process_; } -int Rank::GlobalRank() const { return process_rank_ * thread_size_ + thread_rank_; } +int Rank::GlobalRank() const { return global_process_rank_ * threads_per_process_ + thread_rank_; } -bool Rank::IsParallel() const { return thread_size_ * process_size_ > 1; } +bool Rank::IsParallel() const { return global::GetWorldSize() > 1; } bool Rank::IsMainRank() const { return GlobalRank() == 0; } diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 96776585..f7ced03d 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -7,6 +7,9 @@ include(${CMAKE_SOURCE_DIR}/cmake/test_macros.cmake) # Common test utilities add_subdirectory(common) +# Distributed tests +add_subdirectory(distributed) + # Tensor tests add_subdirectory(tensor) diff --git a/tests/distributed/CMakeLists.txt b/tests/distributed/CMakeLists.txt new file mode 100644 index 00000000..b8ed4970 --- /dev/null +++ b/tests/distributed/CMakeLists.txt @@ -0,0 +1,25 @@ +# ========================================================================== +# Distributed tests +# ========================================================================== + +infini_train_add_test(test_rank + SOURCES test_rank.cc + LABELS cpu +) + +add_test( + NAME RankTest.MultiNodeSingleProcessIsParallel + COMMAND ${CMAKE_COMMAND} -E env + "WORLD_SIZE=2" + "LOCAL_WORLD_SIZE=1" + "RANK=0" + "LOCAL_RANK=0" + $ + --gtest_filter=RankTest.DetectsParallelismFromGlobalWorldSize +) + +set_tests_properties(RankTest.MultiNodeSingleProcessIsParallel + PROPERTIES + LABELS cpu + TIMEOUT 10 +) diff --git a/tests/distributed/test_rank.cc b/tests/distributed/test_rank.cc new file mode 100644 index 00000000..e0b37614 --- /dev/null +++ b/tests/distributed/test_rank.cc @@ -0,0 +1,16 @@ +#include "gtest/gtest.h" + +#include "infini_train/include/nn/parallel/global.h" +#include "infini_train/include/nn/parallel/rank.h" + +namespace infini_train::nn::parallel { +namespace { + +TEST(RankTest, DetectsParallelismFromGlobalWorldSize) { + const Rank rank(global::GetGlobalProcRank(), 0, global::GetNprocPerNode(), global::GetNthreadPerProc()); + + EXPECT_EQ(rank.IsParallel(), global::GetWorldSize() > 1); +} + +} // namespace +} // namespace infini_train::nn::parallel diff --git a/tools/infini_run/infini_run.cc b/tools/infini_run/infini_run.cc index 189d2752..3a32eb6c 100644 --- a/tools/infini_run/infini_run.cc +++ b/tools/infini_run/infini_run.cc @@ -1,4 +1,6 @@ +#include #include +#include #include #include #include @@ -6,6 +8,7 @@ #include #include #include +#include #include #include "gflags/gflags.h" @@ -54,6 +57,20 @@ void SetEnvInt(const char *name, int value) { setenv(name, value_str.c_str(), 1); } +void TerminateChildren(const std::unordered_set &child_pids) { + for (pid_t child_pid : child_pids) { kill(child_pid, SIGTERM); } +} + +int ExitCodeFromStatus(int status) { + if (WIFEXITED(status)) { + return WEXITSTATUS(status); + } + if (WIFSIGNALED(status)) { + return 128 + WTERMSIG(status); + } + return 1; +} + void CleanupRunUniqueIdFiles(const std::string &run_id) { const std::string prefix = "cclUniqueId_" + run_id + "_"; for (const auto &entry : std::filesystem::directory_iterator(std::filesystem::current_path())) { @@ -104,12 +121,19 @@ int main(int argc, char **argv) { std::string master_port = FLAGS_rdzv_endpoint.substr(FLAGS_rdzv_endpoint.find(':') + 1); const std::string run_id = FLAGS_nnodes == 1 ? GenerateLocalRunId() : ""; + std::unordered_set running_children; + int exit_code = 0; + for (int local_proc_rank = 0; local_proc_rank < FLAGS_nproc_per_node; ++local_proc_rank) { pid_t pid = fork(); + if (pid < 0) { + perror("fork failed"); + exit_code = 1; + TerminateChildren(running_children); + break; + } if (pid == 0) { int global_proc_rank = FLAGS_node_rank * FLAGS_nproc_per_node + local_proc_rank; - SetEnvInt("NNODES", FLAGS_nnodes); - SetEnvInt("NPROC_PER_NODE", FLAGS_nproc_per_node); SetEnvInt("LOCAL_WORLD_SIZE", FLAGS_nproc_per_node); setenv("MASTER_ADDR", master_addr.c_str(), 1); @@ -118,45 +142,38 @@ int main(int argc, char **argv) { setenv("INFINI_RUN_ID", run_id.c_str(), 1); } - SetEnvInt("GLOBAL_PROC_RANK", global_proc_rank); - SetEnvInt("LOCAL_PROC_RANK", local_proc_rank); SetEnvInt("RANK", global_proc_rank); SetEnvInt("LOCAL_RANK", local_proc_rank); - SetEnvInt("PROC_WORLD_SIZE", proc_world_size); SetEnvInt("WORLD_SIZE", proc_world_size); - SetEnvInt("GROUP_RANK", FLAGS_node_rank); - SetEnvInt("ROLE_RANK", global_proc_rank); - SetEnvInt("ROLE_WORLD_SIZE", proc_world_size); execvp(train_program.c_str(), train_argv.data()); perror("exec failed"); exit(1); } + running_children.insert(pid); } - int exit_code = 0; - for (int i = 0; i < FLAGS_nproc_per_node; ++i) { + while (!running_children.empty()) { int status; pid_t child = wait(&status); if (child < 0) { + if (errno == EINTR) { + continue; + } perror("wait failed"); - exit_code = 1; + if (exit_code == 0) { + exit_code = 1; + } + TerminateChildren(running_children); break; } - if (WIFEXITED(status)) { - int child_exit_code = WEXITSTATUS(status); - if (child_exit_code != 0 && exit_code == 0) { - exit_code = child_exit_code; - } - } else if (WIFSIGNALED(status)) { - int signal = WTERMSIG(status); - if (exit_code == 0) { - exit_code = 128 + signal; - } - } else if (exit_code == 0) { - exit_code = 1; + running_children.erase(child); + const int child_exit_code = ExitCodeFromStatus(status); + if (child_exit_code != 0 && exit_code == 0) { + exit_code = child_exit_code; + TerminateChildren(running_children); } } From 21dd48a878122ced9b1aaae6348c36819516918b Mon Sep 17 00:00:00 2001 From: cx Date: Thu, 30 Jul 2026 06:06:37 +0000 Subject: [PATCH 5/8] fix: require a shared run ID for multi-node launch - add torchrun-style --rdzv_id support - use the shared ID to isolate CCL unique-ID files - preserve automatic run ID generation for single-node runs - document rdzv_id in the multi-node example --- README.md | 1 + tools/infini_run/infini_run.cc | 15 +++++++++------ 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index befab7a0..ed6f0dd9 100644 --- a/README.md +++ b/README.md @@ -123,6 +123,7 @@ The following examples demonstrate **LLaMA 3 supervised fine-tuning (SFT)** usin --nproc_per_node=1 \ --node_rank=[rank_id] \ --rdzv_endpoint=[master_addr]:29500 \ + --rdzv_id=[job_id] \ ./llama3 \ --device cuda \ --input_bin [training_data_path] \ diff --git a/tools/infini_run/infini_run.cc b/tools/infini_run/infini_run.cc index 3a32eb6c..69f8f0dc 100644 --- a/tools/infini_run/infini_run.cc +++ b/tools/infini_run/infini_run.cc @@ -19,10 +19,13 @@ DEFINE_int32(nproc_per_node, 1, "Number of processes per node"); DEFINE_int32(node_rank, 0, "Rank of this node"); DEFINE_string(rdzv_endpoint, "127.0.0.1:29500", "Rendezvous endpoint (host:port)"); +DEFINE_string(rdzv_id, "", "Unique job ID shared by all nodes"); + namespace { bool IsLauncherFlag(const std::string &flag) { - return flag == "--nnodes" || flag == "--nproc_per_node" || flag == "--node_rank" || flag == "--rdzv_endpoint"; + return flag == "--nnodes" || flag == "--nproc_per_node" || flag == "--node_rank" || flag == "--rdzv_endpoint" + || flag == "--rdzv_id"; } int FindTrainProgramIndex(int argc, char **argv) { @@ -106,6 +109,8 @@ int main(int argc, char **argv) { CHECK_GT(FLAGS_nnodes, 0) << "nnodes must be positive"; CHECK_GT(FLAGS_nproc_per_node, 0) << "nproc_per_node must be positive"; CHECK_NE(FLAGS_rdzv_endpoint.find(':'), std::string::npos) << "rdzv_endpoint must be host:port"; + CHECK(FLAGS_nnodes == 1 || !FLAGS_rdzv_id.empty()) + << "rdzv_id must be set to the same unique job ID on every node for multi-node training"; CHECK_LT(train_program_index, argc) << "No training program specified!"; @@ -119,7 +124,7 @@ int main(int argc, char **argv) { int proc_world_size = FLAGS_nnodes * FLAGS_nproc_per_node; std::string master_addr = FLAGS_rdzv_endpoint.substr(0, FLAGS_rdzv_endpoint.find(':')); std::string master_port = FLAGS_rdzv_endpoint.substr(FLAGS_rdzv_endpoint.find(':') + 1); - const std::string run_id = FLAGS_nnodes == 1 ? GenerateLocalRunId() : ""; + const std::string run_id = FLAGS_rdzv_id.empty() ? GenerateLocalRunId() : FLAGS_rdzv_id; std::unordered_set running_children; int exit_code = 0; @@ -138,9 +143,7 @@ int main(int argc, char **argv) { setenv("MASTER_ADDR", master_addr.c_str(), 1); setenv("MASTER_PORT", master_port.c_str(), 1); - if (!run_id.empty()) { - setenv("INFINI_RUN_ID", run_id.c_str(), 1); - } + setenv("INFINI_RUN_ID", run_id.c_str(), 1); SetEnvInt("RANK", global_proc_rank); SetEnvInt("LOCAL_RANK", local_proc_rank); @@ -177,7 +180,7 @@ int main(int argc, char **argv) { } } - if (!run_id.empty()) { + if (FLAGS_nnodes == 1) { CleanupRunUniqueIdFiles(run_id); } return exit_code; From a1e8738ac972df4817299228b52bf80f5a9380d9 Mon Sep 17 00:00:00 2001 From: cx Date: Thu, 30 Jul 2026 09:41:18 +0000 Subject: [PATCH 6/8] fix: separate launcher args from model test args --- scripts/run_models_and_profile.bash | 2 +- scripts/test_config.json | 13 +++---------- 2 files changed, 4 insertions(+), 11 deletions(-) diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index f3174ab1..1f1e6f05 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -522,7 +522,7 @@ for ((id=0; id Date: Mon, 10 Aug 2026 10:12:26 +0000 Subject: [PATCH 7/8] fix: refine multi-process launcher integration - support infini_run with or without the optional -- separator - validate node rank bounds - use infini_run only for the new 8_proc test group - standardize torchrun environment variables and device index mapping - clarify NCCL unique ID filename helpers --- example/gpt2/main.cc | 2 +- example/llama3/main.cc | 2 +- infini_train/include/nn/parallel/global.h | 2 +- infini_train/src/core/ccl/ccl_utils.cc | 14 +++++------ infini_train/src/device.cc | 11 +++----- infini_train/src/nn/parallel/data_parallel.cc | 3 +-- .../parallel/ddp/distributed_data_parallel.cc | 4 +-- infini_train/src/nn/parallel/global.cc | 25 ++++++++----------- infini_train/src/nn/parallel/process_group.cc | 2 +- scripts/run_models_and_profile.bash | 23 ++++++++++++----- scripts/test_config.json | 2 +- tools/infini_run/infini_run.cc | 6 ++--- 12 files changed, 49 insertions(+), 47 deletions(-) diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 7563c7e2..3fef5bdf 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -180,7 +180,7 @@ void Train(const nn::parallel::Rank &rank) { const ProcessGroup *pp_pg = nullptr; if (rank.IsParallel()) { - device = Device(Device::DeviceType::kCUDA, global::GetLocalDeviceIndex(rank.thread_rank())); + device = Device(Device::DeviceType::kCUDA, global::GetDeviceIndex(rank.thread_rank())); auto *pg_factory = ProcessGroupFactory::Instance(device.type()); if (ddp_world_size > 1) { diff --git a/example/llama3/main.cc b/example/llama3/main.cc index 2510a8f5..ca627741 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -168,7 +168,7 @@ void Train(const nn::parallel::Rank &rank) { const ProcessGroup *pp_pg = nullptr; if (rank.IsParallel()) { - device = Device(Device::DeviceType::kCUDA, global::GetLocalDeviceIndex(rank.thread_rank())); + device = Device(Device::DeviceType::kCUDA, global::GetDeviceIndex(rank.thread_rank())); auto *pg_factory = ProcessGroupFactory::Instance(device.type()); if (ddp_world_size > 1) { diff --git a/infini_train/include/nn/parallel/global.h b/infini_train/include/nn/parallel/global.h index c7c74a99..38694a91 100644 --- a/infini_train/include/nn/parallel/global.h +++ b/infini_train/include/nn/parallel/global.h @@ -98,7 +98,7 @@ inline int GetNprocPerNode() { return GlobalEnv::Instance().nproc_per_node(); } inline int GetNthreadPerProc() { return GlobalEnv::Instance().nthread_per_process(); } inline int GetGlobalProcRank() { return GlobalEnv::Instance().global_proc_rank(); } inline int GetLocalProcRank() { return GlobalEnv::Instance().local_proc_rank(); } -inline int GetLocalDeviceIndex(int thread_rank = 0) { return GetLocalProcRank() * GetNthreadPerProc() + thread_rank; } +inline int GetDeviceIndex(int thread_rank) { return GetLocalProcRank() * GetNthreadPerProc() + thread_rank; } inline int GetTensorParallelSize() { return GlobalEnv::Instance().tensor_parallel_size(); } inline int GetSequenceParallelSize() { return GlobalEnv::Instance().sequence_parallel_size(); } diff --git a/infini_train/src/core/ccl/ccl_utils.cc b/infini_train/src/core/ccl/ccl_utils.cc index ebf1ed19..1368993a 100644 --- a/infini_train/src/core/ccl/ccl_utils.cc +++ b/infini_train/src/core/ccl/ccl_utils.cc @@ -12,13 +12,13 @@ namespace infini_train::core { namespace { -std::string UniqueIdPath(const std::string &pg_name) { +std::string UniqueIdFileName(const std::string &pg_name) { const char *run_id = std::getenv("INFINI_RUN_ID"); const std::string prefix = run_id == nullptr ? "" : std::string(run_id) + "_"; return "cclUniqueId_" + prefix + pg_name + ".bin"; } -std::string UniqueIdTmpPath(const std::string &pg_name) { +std::string UniqueIdTmpFileName(const std::string &pg_name) { const char *run_id = std::getenv("INFINI_RUN_ID"); const std::string prefix = run_id == nullptr ? "" : std::string(run_id) + "_"; return "cclUniqueId_" + prefix + pg_name + ".tmp"; @@ -26,7 +26,7 @@ std::string UniqueIdTmpPath(const std::string &pg_name) { } // namespace void WriteUniqueIdFile(const CclUniqueId &unique_id, const std::string &pg_name) { - const std::string tmp_path = UniqueIdTmpPath(pg_name); + const std::string tmp_path = UniqueIdTmpFileName(pg_name); std::ofstream ofs(tmp_path, std::ios::binary); CHECK(ofs.good()) << "Failed to open unique_id tmp file for write: " << tmp_path; @@ -34,14 +34,14 @@ void WriteUniqueIdFile(const CclUniqueId &unique_id, const std::string &pg_name) ofs.write(reinterpret_cast(unique_id.Data()), static_cast(size)); ofs.close(); - const std::string file_path = UniqueIdPath(pg_name); + const std::string file_path = UniqueIdFileName(pg_name); CHECK_EQ(std::rename(tmp_path.c_str(), file_path.c_str()), 0) << "Failed to rename unique_id file from " << tmp_path << " to " << file_path; } void ReadUniqueIdFile(CclUniqueId *unique_id, const std::string &pg_name) { CHECK_NOTNULL(unique_id); - const std::string file_path = UniqueIdPath(pg_name); + const std::string file_path = UniqueIdFileName(pg_name); while (!std::filesystem::exists(file_path)) { std::this_thread::sleep_for(std::chrono::microseconds(1000)); } @@ -57,12 +57,12 @@ void ReadUniqueIdFile(CclUniqueId *unique_id, const std::string &pg_name) { } void CleanupUniqueIdFile(const std::string &pg_name) { - const std::string file_path = UniqueIdPath(pg_name); + const std::string file_path = UniqueIdFileName(pg_name); if (std::filesystem::exists(file_path)) { std::filesystem::remove(file_path); } - const std::string tmp_path = UniqueIdTmpPath(pg_name); + const std::string tmp_path = UniqueIdTmpFileName(pg_name); if (std::filesystem::exists(tmp_path)) { std::filesystem::remove(tmp_path); } diff --git a/infini_train/src/device.cc b/infini_train/src/device.cc index db10cd54..98021db5 100644 --- a/infini_train/src/device.cc +++ b/infini_train/src/device.cc @@ -33,15 +33,10 @@ std::string Device::ToString() const { } nn::parallel::Rank Device::Rank() const { - if (IsCPU()) { - return {nn::parallel::global::GetGlobalProcRank(), 0, nn::parallel::global::GetNprocPerNode(), - nn::parallel::global::GetNthreadPerProc()}; - } - - const int thread_rank = index_ - nn::parallel::global::GetLocalDeviceIndex(); - CHECK_GE(thread_rank, 0) << "CUDA device index is outside the current process rank range"; + const int thread_rank = index_ - nn::parallel::global::GetDeviceIndex(0); + CHECK_GE(thread_rank, 0) << "Device index is outside the current process rank range"; CHECK_LT(thread_rank, nn::parallel::global::GetNthreadPerProc()) - << "CUDA device index is outside the current process rank range"; + << "Device index is outside the current process rank range"; return {nn::parallel::global::GetGlobalProcRank(), thread_rank, nn::parallel::global::GetNprocPerNode(), nn::parallel::global::GetNthreadPerProc()}; } diff --git a/infini_train/src/nn/parallel/data_parallel.cc b/infini_train/src/nn/parallel/data_parallel.cc index 6199d39f..cf16f6c3 100644 --- a/infini_train/src/nn/parallel/data_parallel.cc +++ b/infini_train/src/nn/parallel/data_parallel.cc @@ -61,8 +61,7 @@ ParallelApply(const std::vector> &modules, DataParallel::DataParallel(const std::shared_ptr &module, int dim, Device::DeviceType device_type) : dim_(dim) { devices_.reserve(global::GetNthreadPerProc()); for (int thread_rank = 0; thread_rank < global::GetNthreadPerProc(); ++thread_rank) { - const int device_index - = device_type == Device::DeviceType::kCUDA ? global::GetLocalDeviceIndex(thread_rank) : thread_rank; + const int device_index = global::GetDeviceIndex(thread_rank); devices_.emplace_back(device_type, device_index); } diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index 0bdf4046..57460ba0 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -33,7 +33,7 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod continue; } auto device = param->GetDevice(); - CHECK_EQ(device.index(), global::GetLocalDeviceIndex(rank.thread_rank())) + CHECK_EQ(device.index(), global::GetDeviceIndex(rank.thread_rank())) << "All parameters must be on the same device as the module"; if (!ddp_config.gradient_bucketing_enabled && ddp_config.zero_stage < 1) { auto hook = std::make_unique( @@ -42,7 +42,7 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod } } for (auto &buffer : module->Buffers()) { - CHECK_EQ(buffer->GetDevice().index(), global::GetLocalDeviceIndex(rank.thread_rank())) + CHECK_EQ(buffer->GetDevice().index(), global::GetDeviceIndex(rank.thread_rank())) << "All buffers must be on the same device as the module"; } modules_[kModuleName] = std::move(module); diff --git a/infini_train/src/nn/parallel/global.cc b/infini_train/src/nn/parallel/global.cc index e56cac72..655b4bce 100644 --- a/infini_train/src/nn/parallel/global.cc +++ b/infini_train/src/nn/parallel/global.cc @@ -92,21 +92,18 @@ void GlobalEnv::Init(int nthread_per_process, int tensor_parallel_size, bool seq CHECK(!initialized_) << "Repeated initialization of GlobalEnv!"; - const int proc_world_size = GetEnvAsInt("WORLD_SIZE", GetEnvAsInt("PROC_WORLD_SIZE", 1)); - nproc_per_node_ = GetEnvAsInt("LOCAL_WORLD_SIZE", GetEnvAsInt("NPROC_PER_NODE", 1)); - CHECK_GT(nproc_per_node_, 0) << "NPROC_PER_NODE/LOCAL_WORLD_SIZE must be positive"; - CHECK_GT(proc_world_size, 0) << "PROC_WORLD_SIZE/WORLD_SIZE must be positive"; - CHECK_EQ(proc_world_size % nproc_per_node_, 0) - << "PROC_WORLD_SIZE/WORLD_SIZE must be divisible by NPROC_PER_NODE/LOCAL_WORLD_SIZE"; + const int proc_world_size = GetEnvAsInt("WORLD_SIZE", 1); + nproc_per_node_ = GetEnvAsInt("LOCAL_WORLD_SIZE", 1); + CHECK_GT(nproc_per_node_, 0) << "LOCAL_WORLD_SIZE must be positive"; + CHECK_GT(proc_world_size, 0) << "WORLD_SIZE must be positive"; + CHECK_EQ(proc_world_size % nproc_per_node_, 0) << "WORLD_SIZE must be divisible by LOCAL_WORLD_SIZE"; nnodes_ = proc_world_size / nproc_per_node_; - global_proc_rank_ = GetEnvAsInt("RANK", GetEnvAsInt("GLOBAL_PROC_RANK", 0)); - local_proc_rank_ = GetEnvAsInt("LOCAL_RANK", GetEnvAsInt("LOCAL_PROC_RANK", 0)); - CHECK_GE(global_proc_rank_, 0) << "GLOBAL_PROC_RANK/RANK must be non-negative"; - CHECK_LT(global_proc_rank_, proc_world_size) - << "GLOBAL_PROC_RANK/RANK must be less than PROC_WORLD_SIZE/WORLD_SIZE"; - CHECK_GE(local_proc_rank_, 0) << "LOCAL_PROC_RANK/LOCAL_RANK must be non-negative"; - CHECK_LT(local_proc_rank_, nproc_per_node_) - << "LOCAL_PROC_RANK/LOCAL_RANK must be less than NPROC_PER_NODE/LOCAL_WORLD_SIZE"; + global_proc_rank_ = GetEnvAsInt("RANK", 0); + local_proc_rank_ = GetEnvAsInt("LOCAL_RANK", 0); + CHECK_GE(global_proc_rank_, 0) << "RANK must be non-negative"; + CHECK_LT(global_proc_rank_, proc_world_size) << "RANK must be less than WORLD_SIZE"; + CHECK_GE(local_proc_rank_, 0) << "LOCAL_RANK must be non-negative"; + CHECK_LT(local_proc_rank_, nproc_per_node_) << "LOCAL_RANK must be less than LOCAL_WORLD_SIZE"; nthread_per_process_ = nthread_per_process; world_size_ = proc_world_size * nthread_per_process; diff --git a/infini_train/src/nn/parallel/process_group.cc b/infini_train/src/nn/parallel/process_group.cc index fda07e13..caa08981 100644 --- a/infini_train/src/nn/parallel/process_group.cc +++ b/infini_train/src/nn/parallel/process_group.cc @@ -99,7 +99,7 @@ void ProcessGroup::InitMultiProcess(const std::vector &ranks) { int global_thread_rank = lower_rank + i; auto it = std::ranges::find(ranks, global_thread_rank); if (it != ranks.end()) { - auto device = Device(backend_, global::GetLocalDeviceIndex(i)); + auto device = Device(backend_, global::GetDeviceIndex(i)); core::DeviceGuard guard(device); core::CclComm *comm_raw = nullptr; diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index 1f1e6f05..055ab47d 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -439,7 +439,7 @@ check_model_inputs() { fi } -model_cmd_for_test() { +infini_run_cmd_for_test() { local model_bin="$1" local input_bin="$2" local llmc_filepath="$3" @@ -522,25 +522,36 @@ for ((id=0; id train_argv; for (int i = train_program_index; i < argc; ++i) { train_argv.push_back(argv[i]); } train_argv.push_back(nullptr); From ab6ac04ee2c2acfbc2efc1047523e0fc9779d753 Mon Sep 17 00:00:00 2001 From: cx Date: Tue, 11 Aug 2026 06:27:22 +0000 Subject: [PATCH 8/8] fix(scripts): rename multi-process test cases - add an _8_proc suffix to multi-process test case IDs - print manual comparison commands when no baseline log directory is set --- scripts/run_models_and_profile.bash | 4 ++++ scripts/test_config.json | 20 ++++++++++---------- 2 files changed, 14 insertions(+), 10 deletions(-) diff --git a/scripts/run_models_and_profile.bash b/scripts/run_models_and_profile.bash index 055ab47d..67bc7acd 100755 --- a/scripts/run_models_and_profile.bash +++ b/scripts/run_models_and_profile.bash @@ -587,6 +587,10 @@ else echo -e "\033[1;33m To enable comparison, set 'variables.COMPARE_LOG_DIR' in ${CONFIG_FILE}\033[0m" echo -e "\033[1;33m or export COMPARE_LOG_DIR=/path/to/baseline_logs before running.\033[0m" echo -e "\033[1;33m============================================================\033[0m" + + echo -e "\n\033[1;33mRun comparison manually:\033[0m" + echo "python3 compare_loss.py \"/path/to/baseline_logs\" \"$(realpath "$LOG_DIR")\"" + echo "python3 compare_tps.py \"/path/to/baseline_logs\" \"$(realpath "$LOG_DIR")\"" fi echo -e "\n\033[1;36m[END OF TEST] Cleaning build directory after all tests\033[0m" diff --git a/scripts/test_config.json b/scripts/test_config.json index 65e10c64..2b8202d2 100644 --- a/scripts/test_config.json +++ b/scripts/test_config.json @@ -903,7 +903,7 @@ }, "tests": [ { - "id": "3", + "id": "3_8proc", "args": { "dtype": "float32", "nthread_per_process": 1, @@ -913,7 +913,7 @@ } }, { - "id": "3_bfloat16", + "id": "3_bfloat16_8proc", "args": { "dtype": "bfloat16", "nthread_per_process": 1, @@ -923,7 +923,7 @@ } }, { - "id": "4", + "id": "4_8proc", "args": { "dtype": "float32", "nthread_per_process": 1, @@ -934,7 +934,7 @@ } }, { - "id": "4_bfloat16", + "id": "4_bfloat16_8proc", "args": { "dtype": "bfloat16", "nthread_per_process": 1, @@ -945,7 +945,7 @@ } }, { - "id": "5", + "id": "5_8proc", "args": { "dtype": "float32", "nthread_per_process": 1, @@ -957,7 +957,7 @@ } }, { - "id": "5_bfloat16", + "id": "5_bfloat16_8proc", "args": { "dtype": "bfloat16", "nthread_per_process": 1, @@ -969,7 +969,7 @@ } }, { - "id": "6", + "id": "6_8proc", "args": { "dtype": "float32", "nthread_per_process": 1, @@ -980,7 +980,7 @@ } }, { - "id": "6_bfloat16", + "id": "6_bfloat16_8proc", "args": { "dtype": "bfloat16", "nthread_per_process": 1, @@ -991,7 +991,7 @@ } }, { - "id": "8", + "id": "8_8proc", "args": { "dtype": "float32", "nthread_per_process": 1, @@ -1005,7 +1005,7 @@ } }, { - "id": "8_bfloat16", + "id": "8_bfloat16_8proc", "args": { "dtype": "bfloat16", "nthread_per_process": 1,