diff --git a/orttraining/orttraining/training_ops/cuda/communication/nccl_service.cc b/orttraining/orttraining/training_ops/cuda/communication/nccl_service.cc index 449b8fee2b..7118affb1e 100644 --- a/orttraining/orttraining/training_ops/cuda/communication/nccl_service.cc +++ b/orttraining/orttraining/training_ops/cuda/communication/nccl_service.cc @@ -229,21 +229,12 @@ void NcclService::Initialize() { // CPUs // Other devices - // Make sure MPI is initialized. - int is_mpi_initialized = 0; - MPI_CHECK(MPI_Initialized(&is_mpi_initialized)); - if (!is_mpi_initialized) { - int mpi_threads_provided = 0; - MPI_CHECK(MPI_Init_thread(nullptr, nullptr, MPI_THREAD_MULTIPLE, &mpi_threads_provided)); - } - - int mpi_rank; - int mpi_size; - MPI_CHECK(MPI_Comm_rank(MPI_COMM_WORLD, &mpi_rank)); - MPI_CHECK(MPI_Comm_size(MPI_COMM_WORLD, &mpi_size)); + const int mpi_rank = onnxruntime::training::MPIContext::GetInstance().GetWorldRank(); + const int mpi_local_rank = onnxruntime::training::MPIContext::GetInstance().GetLocalRank(); + const int mpi_size = onnxruntime::training::MPIContext::GetInstance().GetWorldSize(); // Set device this NCCL communicator runs on. - CUDA_CALL(cudaSetDevice(mpi_rank)); + CUDA_CALL(cudaSetDevice(mpi_local_rank)); // Create communication stream. CUDA_CALL(cudaStreamCreateWithFlags(&stream_, cudaStreamNonBlocking));