From a2b0101ec43a227fab598632a0467d69943ddb02 Mon Sep 17 00:00:00 2001 From: Li Baoming <1508269885@qq.com> Date: Wed, 2 Sep 2026 11:12:54 +0800 Subject: [PATCH 1/7] feat(cambricon): add CNCL backend for local multi-device collectives --- CMakeLists.txt | 27 +++++++- scripts/gen_bridge.py | 3 + src/CMakeLists.txt | 10 +++ src/backend.h | 5 ++ src/backend_device_map.h | 4 ++ src/backends/ccl/cncl/api.h | 58 +++++++++++++++++ src/backends/ccl/cncl/cambricon/api.h | 15 +++++ src/backends/ccl/cncl/impl/all_gather.h | 17 +++++ src/backends/ccl/cncl/impl/all_reduce.h | 17 +++++ src/backends/ccl/cncl/impl/comm_destroy.h | 17 +++++ src/backends/ccl/cncl/impl/comm_init_all.h | 75 ++++++++++++++++++++++ src/backends/ccl/cncl/type_map.h | 73 +++++++++++++++++++++ src/backends/mpi/ompi/impl/comm_init_all.h | 17 +++-- src/base/comm_init_all.h | 21 ++---- 14 files changed, 339 insertions(+), 20 deletions(-) create mode 100644 src/backends/ccl/cncl/api.h create mode 100644 src/backends/ccl/cncl/cambricon/api.h create mode 100644 src/backends/ccl/cncl/impl/all_gather.h create mode 100644 src/backends/ccl/cncl/impl/all_reduce.h create mode 100644 src/backends/ccl/cncl/impl/comm_destroy.h create mode 100644 src/backends/ccl/cncl/impl/comm_init_all.h create mode 100644 src/backends/ccl/cncl/type_map.h diff --git a/CMakeLists.txt b/CMakeLists.txt index 1b006ce..48181fc 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -23,6 +23,7 @@ option(WITH_OMPI "Enable OpenMPI backend" OFF) option(WITH_MPICH "Enable MPICH backend" OFF) option(WITH_NCCL "Enable NCCL backend" OFF) option(WITH_MCCL "Enable MCCL backend" OFF) +option(WITH_CNCL "Enable CNCL backend" OFF) # ========================================================= # --- MISC. BUILD OPTIONS --- @@ -327,10 +328,23 @@ if(AUTO_DETECT_BACKENDS) else() message(STATUS "No suitable device environment, skipping MCCL detection.") endif() + + # Detect CNCL Dependencies + if(WITH_CAMBRICON) + find_path(AUTO_CNCL_INC NAMES cncl.h HINTS "$ENV{NEUWARE_HOME}" /usr/local/neuware PATH_SUFFIXES include QUIET) + find_library(AUTO_CNCL_LIB NAMES cncl HINTS "$ENV{NEUWARE_HOME}" /usr/local/neuware PATH_SUFFIXES lib lib64 QUIET) + + if(AUTO_CNCL_INC AND AUTO_CNCL_LIB) + set(WITH_CNCL ON) + message(STATUS "Auto-detected CNCL backend.") + else() + message(STATUS "CNCL library/headers not found in Cambricon paths.") + endif() + endif() endif() # Fallback: If no backends are enabled or auto-detected, fall back to OpenMPI as the default bootstrap profile. -if(NOT WITH_OMPI AND NOT WITH_MPICH AND NOT WITH_NCCL AND NOT WITH_MCCL) +if(NOT WITH_OMPI AND NOT WITH_MPICH AND NOT WITH_NCCL AND NOT WITH_MCCL AND NOT WITH_CNCL) set(WITH_OMPI ON) message(STATUS "No backend specified or detected. Defaulting to `WITH_OMPI=ON`") endif() @@ -535,6 +549,17 @@ if(WITH_MCCL) include_directories(${MCCL_INC}) endif() +if(WITH_CNCL) + if(NOT WITH_CAMBRICON) + message(FATAL_ERROR "CNCL backend requires Cambricon device support. Please enable `WITH_CAMBRICON`.") + endif() + + find_library(CNCL_LIB NAMES cncl HINTS "${NEUWARE_HOME}" "$ENV{NEUWARE_HOME}" /usr/local/neuware PATH_SUFFIXES lib lib64 REQUIRED) + find_path(CNCL_INC NAMES cncl.h HINTS "${NEUWARE_HOME}" "$ENV{NEUWARE_HOME}" /usr/local/neuware PATH_SUFFIXES include REQUIRED) + + include_directories(${CNCL_INC}) +endif() + # Python is required for code generation. find_package(Python3 REQUIRED) diff --git a/scripts/gen_bridge.py b/scripts/gen_bridge.py index 211080b..59d4dd2 100644 --- a/scripts/gen_bridge.py +++ b/scripts/gen_bridge.py @@ -32,16 +32,19 @@ "mpich": ["backends/mpi/ompi/impl"], "nccl": ["backends/ccl/nccl/impl"], "mccl": ["backends/ccl/mccl/impl"], + "cncl": ["backends/ccl/cncl/impl"], } BACKEND_COMMON_HEADERS = { "nccl": ["backends/ccl/nccl/type_map.h"], "mccl": ["backends/ccl/mccl/type_map.h"], + "cncl": ["backends/ccl/cncl/type_map.h"], } CCL_PROVIDER_BACKENDS = { "nccl": "backends/ccl/nccl", "mccl": "backends/ccl/mccl", + "cncl": "backends/ccl/cncl", } # ================================================================= diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 84ffbcf..c191b68 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -283,6 +283,16 @@ if(WITH_MCCL) target_link_libraries(infiniccl PRIVATE ${MCCL_LIB}) endif() +# CNCL +if(WITH_CNCL) + list(APPEND BACKEND_LIST "cncl") + file(GLOB_RECURSE CNCL_SRCS "backends/ccl/cncl/*.cc" "backends/ccl/cncl/*.cpp") + + target_sources(infiniccl PRIVATE ${CNCL_SRCS}) + target_include_directories(infiniccl PRIVATE ${CNCL_INC}) + target_link_libraries(infiniccl PRIVATE ${CNCL_LIB}) +endif() + # ========================================================= # --- File Generation --- # ========================================================= diff --git a/src/backend.h b/src/backend.h index 0b02662..7ebad6d 100644 --- a/src/backend.h +++ b/src/backend.h @@ -63,6 +63,11 @@ struct BackendPriority { static constexpr int value = 10; }; +template <> +struct BackendPriority { + static constexpr int value = 10; +}; + } // namespace infini::ccl #endif // INFINI_CCL_BACKEND_H_ diff --git a/src/backend_device_map.h b/src/backend_device_map.h index 7cf46b0..8c65d82 100644 --- a/src/backend_device_map.h +++ b/src/backend_device_map.h @@ -33,6 +33,10 @@ template <> struct IsSupportedCombination : std::true_type {}; +template <> +struct IsSupportedCombination + : std::true_type {}; + }; // namespace infini::ccl #endif // INFINI_CCL_BACKEND_DEVICE_MAP_H_ diff --git a/src/backends/ccl/cncl/api.h b/src/backends/ccl/cncl/api.h new file mode 100644 index 0000000..709da32 --- /dev/null +++ b/src/backends/ccl/cncl/api.h @@ -0,0 +1,58 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_API_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_API_H_ + +#include + +#include + +#include "backends/ccl/common/api.h" +#include "logging.h" +#include "return_status_impl.h" +#include "runtime.h" + +namespace infini::ccl { + +template +struct CnclApi { + static constexpr BackendType kBackendType = BackendType::kCncl; + static constexpr Device::Type kDeviceType = device; + + using Comm = cnclComm_t; + using Result = cnclResult_t; + using DataType = cnclDataType_t; + using RedOp = cnclReduceOp_t; + using Stream = typename Runtime::Stream; + + static ReturnStatus Check(Result result) { + if (result != CNCL_RET_SUCCESS) { + LOG(cnclGetErrorStr(result)); + return ReturnStatus::kSystemError; + } + return ReturnStatus::kSuccess; + } + + static Result CommInitAll(Comm* comms, int n_dev, const int* dev_list, + const int* rank_list) { + return cnclInitComms(comms, n_dev, dev_list, rank_list, n_dev, nullptr); + } + + static Result CommDestroy(Comm comm) { return cnclFreeComm(comm); } + + static Result AllReduce(const void* send_buff, void* recv_buff, size_t count, + DataType data_type, RedOp op, Comm comm, + Stream stream) { + return cnclAllReduce(send_buff, recv_buff, count, data_type, op, comm, + stream); + } + + static Result AllGather(const void* send_buff, void* recv_buff, + size_t send_count, DataType data_type, Comm comm, + Stream stream) { + return cnclAllGather(send_buff, recv_buff, send_count, data_type, comm, + stream); + } +}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_API_H_ diff --git a/src/backends/ccl/cncl/cambricon/api.h b/src/backends/ccl/cncl/cambricon/api.h new file mode 100644 index 0000000..865fc0e --- /dev/null +++ b/src/backends/ccl/cncl/cambricon/api.h @@ -0,0 +1,15 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_CAMBRICON_API_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_CAMBRICON_API_H_ + +#include "backends/ccl/cncl/api.h" +#include "devices/cambricon/runtime_.h" + +namespace infini::ccl { + +template <> +struct CclApi + : CnclApi {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_CAMBRICON_API_H_ diff --git a/src/backends/ccl/cncl/impl/all_gather.h b/src/backends/ccl/cncl/impl/all_gather.h new file mode 100644 index 0000000..6ab1dad --- /dev/null +++ b/src/backends/ccl/cncl/impl/all_gather.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_GATHER_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_GATHER_H_ + +#include "backends/ccl/common/impl/all_gather.h" + +namespace infini::ccl { + +template +class AllGatherImpl + : public CclAllGatherImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_GATHER_H_ diff --git a/src/backends/ccl/cncl/impl/all_reduce.h b/src/backends/ccl/cncl/impl/all_reduce.h new file mode 100644 index 0000000..490f609 --- /dev/null +++ b/src/backends/ccl/cncl/impl/all_reduce.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_REDUCE_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_REDUCE_H_ + +#include "backends/ccl/common/impl/all_reduce.h" + +namespace infini::ccl { + +template +class AllReduceImpl + : public CclAllReduceImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_ALL_REDUCE_H_ diff --git a/src/backends/ccl/cncl/impl/comm_destroy.h b/src/backends/ccl/cncl/impl/comm_destroy.h new file mode 100644 index 0000000..efc8899 --- /dev/null +++ b/src/backends/ccl/cncl/impl/comm_destroy.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_DESTROY_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_DESTROY_H_ + +#include "backends/ccl/common/impl/comm_destroy.h" + +namespace infini::ccl { + +template +class CommDestroyImpl + : public CclCommDestroyImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_DESTROY_H_ diff --git a/src/backends/ccl/cncl/impl/comm_init_all.h b/src/backends/ccl/cncl/impl/comm_init_all.h new file mode 100644 index 0000000..06fdfa2 --- /dev/null +++ b/src/backends/ccl/cncl/impl/comm_init_all.h @@ -0,0 +1,75 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ + +#include +#include +#include + +#include "backends/ccl/common/comm_instance.h" +#include "base/comm_init_all.h" +#include "communicator.h" +#include "runtime.h" + +namespace infini::ccl { + +template +class CommInitAllImpl { + public: + static ReturnStatus Apply(void** comm_handles, int n_dev, + const int* dev_list) { + using Api = CclApi; + using CommInstance = CclCommInstance; + using Rt = Runtime; + + if (!comm_handles || !dev_list || n_dev <= 0) { + return ReturnStatus::kInvalidArgument; + } + + auto** comms = reinterpret_cast(comm_handles); + for (int i = 0; i < n_dev; ++i) { + if (comms[i]) { + return ReturnStatus::kInvalidArgument; + } + } + + std::vector> wrappers; + std::vector> instances; + std::vector backend_comms(n_dev); + std::vector rank_list(n_dev); + wrappers.reserve(n_dev); + instances.reserve(n_dev); + std::iota(rank_list.begin(), rank_list.end(), 0); + + for (int i = 0; i < n_dev; ++i) { + auto status = Rt::Check(Rt::SetDevice(dev_list[i])); + if (status != ReturnStatus::kSuccess) { + return status; + } + wrappers.emplace_back( + std::make_unique(device, dev_list[i])); + instances.emplace_back(std::make_unique()); + } + + auto status = Api::Check(Api::CommInitAll(backend_comms.data(), n_dev, + dev_list, rank_list.data())); + if (status != ReturnStatus::kSuccess) { + return status; + } + + for (int i = 0; i < n_dev; ++i) { + instances[i]->handle = backend_comms[i]; + wrappers[i]->set_world_info(i, n_dev); + wrappers[i]->set_intra_comm(std::move(instances[i])); + comms[i] = wrappers[i].release(); + } + + return ReturnStatus::kSuccess; + } +}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ diff --git a/src/backends/ccl/cncl/type_map.h b/src/backends/ccl/cncl/type_map.h new file mode 100644 index 0000000..04f6eee --- /dev/null +++ b/src/backends/ccl/cncl/type_map.h @@ -0,0 +1,73 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_H_ + +#include + +#include "backends/ccl/common/api.h" +#include "comm_impl.h" +#include "data_type_impl.h" + +namespace infini::ccl { + +inline bool DataTypeToCnclType(DataType dtype, cnclDataType_t* cncl_dtype) { + switch (dtype) { + case DataType::kInt8: + *cncl_dtype = cnclInt8; + return true; + case DataType::kInt16: + *cncl_dtype = cnclInt16; + return true; + case DataType::kInt32: + *cncl_dtype = cnclInt32; + return true; + case DataType::kInt64: + *cncl_dtype = cnclInt64; + return true; + case DataType::kUInt8: + *cncl_dtype = cnclUint8; + return true; + case DataType::kUInt16: + *cncl_dtype = cnclUint16; + return true; + case DataType::kUInt32: + *cncl_dtype = cnclUint32; + return true; + case DataType::kUInt64: + *cncl_dtype = cnclUint64; + return true; + case DataType::kFloat16: + *cncl_dtype = cnclFloat16; + return true; + case DataType::kBFloat16: + *cncl_dtype = cnclBfloat16; + return true; + case DataType::kFloat32: + *cncl_dtype = cnclFloat32; + return true; + default: + return false; + } +} + +template <> +struct CclTypeMap { + using Api = CclApi; + + static bool ToBackendDataType(DataType dtype, + typename Api::DataType* backend_dtype) { + return DataTypeToCnclType(dtype, backend_dtype); + } + + static bool ToBackendRedOp(ReductionOpType red_op, + typename Api::RedOp* backend_op) { + if (red_op == ReductionOpType::kAvg) { + return false; + } + *backend_op = static_cast(red_op); + return true; + } +}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_H_ diff --git a/src/backends/mpi/ompi/impl/comm_init_all.h b/src/backends/mpi/ompi/impl/comm_init_all.h index e834622..7cf320b 100644 --- a/src/backends/mpi/ompi/impl/comm_init_all.h +++ b/src/backends/mpi/ompi/impl/comm_init_all.h @@ -12,19 +12,28 @@ namespace infini::ccl { template class CommInitAllImpl { public: - static ReturnStatus Apply(Communicator *comm, int n_dev, - const int *dev_list) { + static ReturnStatus Apply(void** comm_handles, int n_dev, + const int* dev_list) { constexpr Device::Type kDev = ListGetBest(ActiveDevices{}); using Rt = Runtime; - if (!comm) { + if (!comm_handles) { // TODO(lzm): change to use `glog`. LOG("Failed to initialize OpenMPI communicator: invalid " "communicator pointer."); return ReturnStatus::kInternalError; } + Communicator*& comm = *reinterpret_cast(comm_handles); + if (comm && comm->inter_comm()) { + LOG("Invalid communicator handle for `CommInitAll`."); + return ReturnStatus::kInvalidArgument; + } + if (!comm) { + comm = new Communicator(kDev, 0); + } + int rank, size; auto inst = std::make_unique(); INFINI_CHECK_MPI(MPI_Comm_dup(MPI_COMM_WORLD, &inst->handle)); @@ -35,7 +44,7 @@ class CommInitAllImpl { comm->set_inter_comm(std::move(inst)); int local_rank = 0; - char *local_rank_str = getenv("OMPI_COMM_WORLD_LOCAL_RANK"); + char* local_rank_str = getenv("OMPI_COMM_WORLD_LOCAL_RANK"); if (local_rank_str) { local_rank = atoi(local_rank_str); } diff --git a/src/base/comm_init_all.h b/src/base/comm_init_all.h index e737306..2371e50 100644 --- a/src/base/comm_init_all.h +++ b/src/base/comm_init_all.h @@ -13,25 +13,16 @@ struct CommInitAllImpl; class CommInitAll : public Operation { public: - template - static ReturnStatus Execute(void **comm_handle, Args &&...args) { - Communicator *&comm = *reinterpret_cast(comm_handle); - if (comm && comm->inter_comm()) { - // TODO(lzm): change to use `glog`. + template + static ReturnStatus Execute(void** comm_handles, int n_dev, + const int* dev_list) { + if (!comm_handles || n_dev <= 0) { LOG("Invalid communicator handle for `CommInitAll`."); return ReturnStatus::kInvalidArgument; } - constexpr Device::Type kDev = - ListGetBest(ActiveDevices{}); - - if (!comm) { - comm = new Communicator(kDev, 0); - } - - return CommInitAllImpl::Apply( - comm, std::forward(args)...); + return CommInitAllImpl::Apply(comm_handles, + n_dev, dev_list); } }; From 75c7aa4690e08ec8f6874e287f7b29c75c7f11f9 Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Thu, 17 Sep 2026 18:29:20 +0800 Subject: [PATCH 2/7] feat: support the missing CCL implementations for CNCl and orgnaize relevant code - support `GetUniqueId` and `CommInitRank` for CNCL - remove the currently unsupported and irrelevant CCL backend of `CommInitAll` - set `INFINICCL_UNIQUE_ID_BYTES` to 136 to accommodate `cnclCliqueId` - organize code for `CommInitAll` and `CommInitRank` --- include/comm.h | 2 +- src/backends/ccl/cncl/api.h | 15 +++- src/backends/ccl/cncl/impl/comm_init_all.h | 75 ------------------- src/backends/ccl/cncl/impl/comm_init_rank.h | 17 +++++ src/backends/ccl/cncl/impl/get_unique_id.h | 17 +++++ src/backends/ccl/common/impl/comm_init_rank.h | 13 +++- src/backends/mpi/ompi/impl/comm_init_all.h | 18 +---- src/base/comm_init_all.h | 13 +++- src/base/comm_init_rank.h | 10 +-- 9 files changed, 77 insertions(+), 103 deletions(-) delete mode 100644 src/backends/ccl/cncl/impl/comm_init_all.h create mode 100644 src/backends/ccl/cncl/impl/comm_init_rank.h create mode 100644 src/backends/ccl/cncl/impl/get_unique_id.h diff --git a/include/comm.h b/include/comm.h index ec16c92..e36f2ac 100644 --- a/include/comm.h +++ b/include/comm.h @@ -10,7 +10,7 @@ extern "C" { #endif -#define INFINICCL_UNIQUE_ID_BYTES 128 +#define INFINICCL_UNIQUE_ID_BYTES 136 typedef void *infinicclComm_t; diff --git a/src/backends/ccl/cncl/api.h b/src/backends/ccl/cncl/api.h index 709da32..b2b8ca2 100644 --- a/src/backends/ccl/cncl/api.h +++ b/src/backends/ccl/cncl/api.h @@ -18,6 +18,7 @@ struct CnclApi { static constexpr Device::Type kDeviceType = device; using Comm = cnclComm_t; + using UniqueId = cnclCliqueId; using Result = cnclResult_t; using DataType = cnclDataType_t; using RedOp = cnclReduceOp_t; @@ -31,9 +32,17 @@ struct CnclApi { return ReturnStatus::kSuccess; } - static Result CommInitAll(Comm* comms, int n_dev, const int* dev_list, - const int* rank_list) { - return cnclInitComms(comms, n_dev, dev_list, rank_list, n_dev, nullptr); + static Result GetUniqueId(UniqueId* id) { return cnclGetCliqueId(id); } + + static Result CommInitRank(Comm* comm, int nranks, UniqueId id, int rank) { + using Rt = Runtime; + + int device_id = 0; + if (Rt::GetDevice(&device_id) != cnrtSuccess) { + return CNCL_RET_ERR_MLU_RUNTIME; + } + + return cnclInitComms(comm, 1, &device_id, &rank, nranks, &id); } static Result CommDestroy(Comm comm) { return cnclFreeComm(comm); } diff --git a/src/backends/ccl/cncl/impl/comm_init_all.h b/src/backends/ccl/cncl/impl/comm_init_all.h deleted file mode 100644 index 06fdfa2..0000000 --- a/src/backends/ccl/cncl/impl/comm_init_all.h +++ /dev/null @@ -1,75 +0,0 @@ -#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ -#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ - -#include -#include -#include - -#include "backends/ccl/common/comm_instance.h" -#include "base/comm_init_all.h" -#include "communicator.h" -#include "runtime.h" - -namespace infini::ccl { - -template -class CommInitAllImpl { - public: - static ReturnStatus Apply(void** comm_handles, int n_dev, - const int* dev_list) { - using Api = CclApi; - using CommInstance = CclCommInstance; - using Rt = Runtime; - - if (!comm_handles || !dev_list || n_dev <= 0) { - return ReturnStatus::kInvalidArgument; - } - - auto** comms = reinterpret_cast(comm_handles); - for (int i = 0; i < n_dev; ++i) { - if (comms[i]) { - return ReturnStatus::kInvalidArgument; - } - } - - std::vector> wrappers; - std::vector> instances; - std::vector backend_comms(n_dev); - std::vector rank_list(n_dev); - wrappers.reserve(n_dev); - instances.reserve(n_dev); - std::iota(rank_list.begin(), rank_list.end(), 0); - - for (int i = 0; i < n_dev; ++i) { - auto status = Rt::Check(Rt::SetDevice(dev_list[i])); - if (status != ReturnStatus::kSuccess) { - return status; - } - wrappers.emplace_back( - std::make_unique(device, dev_list[i])); - instances.emplace_back(std::make_unique()); - } - - auto status = Api::Check(Api::CommInitAll(backend_comms.data(), n_dev, - dev_list, rank_list.data())); - if (status != ReturnStatus::kSuccess) { - return status; - } - - for (int i = 0; i < n_dev; ++i) { - instances[i]->handle = backend_comms[i]; - wrappers[i]->set_world_info(i, n_dev); - wrappers[i]->set_intra_comm(std::move(instances[i])); - comms[i] = wrappers[i].release(); - } - - return ReturnStatus::kSuccess; - } -}; - -template <> -struct BackendEnabled : std::true_type {}; - -} // namespace infini::ccl - -#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_ALL_H_ diff --git a/src/backends/ccl/cncl/impl/comm_init_rank.h b/src/backends/ccl/cncl/impl/comm_init_rank.h new file mode 100644 index 0000000..8ca3f59 --- /dev/null +++ b/src/backends/ccl/cncl/impl/comm_init_rank.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ + +#include "backends/ccl/common/impl/comm_init_rank.h" + +namespace infini::ccl { + +template +class CommInitRankImpl + : public CclCommInitRankImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ diff --git a/src/backends/ccl/cncl/impl/get_unique_id.h b/src/backends/ccl/cncl/impl/get_unique_id.h new file mode 100644 index 0000000..fac71f2 --- /dev/null +++ b/src/backends/ccl/cncl/impl/get_unique_id.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_GET_UNIQUE_ID_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_GET_UNIQUE_ID_H_ + +#include "backends/ccl/common/impl/get_unique_id.h" + +namespace infini::ccl { + +template +class GetUniqueIdImpl + : public CclGetUniqueIdImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_GET_UNIQUE_ID_H_ diff --git a/src/backends/ccl/common/impl/comm_init_rank.h b/src/backends/ccl/common/impl/comm_init_rank.h index 2b7adee..aced7b5 100644 --- a/src/backends/ccl/common/impl/comm_init_rank.h +++ b/src/backends/ccl/common/impl/comm_init_rank.h @@ -1,6 +1,7 @@ #ifndef INFINI_CCL_BACKENDS_CCL_COMMON_IMPL_COMM_INIT_RANK_H_ #define INFINI_CCL_BACKENDS_CCL_COMMON_IMPL_COMM_INIT_RANK_H_ +#include #include #include "backends/ccl/common/api.h" @@ -19,12 +20,18 @@ class CclCommInitRankImpl { using Api = CclApi; using CommInstance = CclCommInstance; - const auto *backend_id = - reinterpret_cast(id.internal); + if (comm && comm->intra_comm()) { + // TODO(lzm): change to use `glog`. + LOG("Invalid communicator handle for `CommInitRank`."); + return ReturnStatus::kInvalidArgument; + } + + typename Api::UniqueId backend_id{}; + std::memcpy(&backend_id, id.internal, sizeof(backend_id)); typename Api::Comm ccl_handle{}; auto status = - Api::Check(Api::CommInitRank(&ccl_handle, nranks, *backend_id, rank)); + Api::Check(Api::CommInitRank(&ccl_handle, nranks, backend_id, rank)); if (status != ReturnStatus::kSuccess) { return status; } diff --git a/src/backends/mpi/ompi/impl/comm_init_all.h b/src/backends/mpi/ompi/impl/comm_init_all.h index 7cf320b..3217655 100644 --- a/src/backends/mpi/ompi/impl/comm_init_all.h +++ b/src/backends/mpi/ompi/impl/comm_init_all.h @@ -12,27 +12,17 @@ namespace infini::ccl { template class CommInitAllImpl { public: - static ReturnStatus Apply(void** comm_handles, int n_dev, - const int* dev_list) { + static ReturnStatus Apply(Communicator *comm, int n_dev, + const int *dev_list) { constexpr Device::Type kDev = ListGetBest(ActiveDevices{}); using Rt = Runtime; - if (!comm_handles) { - // TODO(lzm): change to use `glog`. - LOG("Failed to initialize OpenMPI communicator: invalid " - "communicator pointer."); - return ReturnStatus::kInternalError; - } - - Communicator*& comm = *reinterpret_cast(comm_handles); if (comm && comm->inter_comm()) { + // TODO(lzm): change to use `glog`. LOG("Invalid communicator handle for `CommInitAll`."); return ReturnStatus::kInvalidArgument; } - if (!comm) { - comm = new Communicator(kDev, 0); - } int rank, size; auto inst = std::make_unique(); @@ -44,7 +34,7 @@ class CommInitAllImpl { comm->set_inter_comm(std::move(inst)); int local_rank = 0; - char* local_rank_str = getenv("OMPI_COMM_WORLD_LOCAL_RANK"); + char *local_rank_str = getenv("OMPI_COMM_WORLD_LOCAL_RANK"); if (local_rank_str) { local_rank = atoi(local_rank_str); } diff --git a/src/base/comm_init_all.h b/src/base/comm_init_all.h index 2371e50..c37d781 100644 --- a/src/base/comm_init_all.h +++ b/src/base/comm_init_all.h @@ -21,8 +21,17 @@ class CommInitAll : public Operation { return ReturnStatus::kInvalidArgument; } - return CommInitAllImpl::Apply(comm_handles, - n_dev, dev_list); + constexpr Device::Type kDev = + ListGetBest(ActiveDevices{}); + + Communicator *&comm = *reinterpret_cast(comm_handle); + + if (!comm) { + comm = new Communicator(kDev, 0); + } + + return CommInitAllImpl::Apply( + comm, std::forward(args)...); } }; diff --git a/src/base/comm_init_rank.h b/src/base/comm_init_rank.h index 46f93f2..e4391d4 100644 --- a/src/base/comm_init_rank.h +++ b/src/base/comm_init_rank.h @@ -14,11 +14,9 @@ class CommInitRank : public Operation { public: template - static ReturnStatus Execute(void **comm_handle, Args &&...args) { - Communicator *&comm = *reinterpret_cast(comm_handle); - if (comm && comm->intra_comm()) { - // TODO(lzm): change to use `glog`. - LOG("Invalid communicator handle for `CommInitRank`."); + static ReturnStatus Execute(void **comm_handle, Args &&...args) { + if (!comm_handle) { + LOG("Invalid communicator handle for `CommInitAll`."); return ReturnStatus::kInvalidArgument; } @@ -29,6 +27,8 @@ class CommInitRank : public Operation { int current_dev = 0; CHECK_STATUS(Rt, Rt::GetDevice(¤t_dev)); + Communicator *&comm = *reinterpret_cast(comm_handle); + if (!comm) { comm = new Communicator(kDev, current_dev); } else { From 473232bebb66ffe70faa23190d7ba290006f56ef Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Wed, 23 Sep 2026 16:50:57 +0800 Subject: [PATCH 3/7] feat: support CNCL rank and point-to-point communication - coordinate process-local `CommInitRank` requests into one `cnclInitComms` call per clique - support native multi-threaded, MPI-hybrid, and MxN initialization - validate local ranks, global ranks, devices, and deferred initialization groups - add CNCL `Send` and `Recv` providers through the common CCL implementation - provide a fallback queue for CNCL versions that reject `nullptr` queues - add CNCL error checking and preserve local communicator metadata --- src/backends/ccl/cncl/api.h | 80 ++++++- src/backends/ccl/cncl/checks.h | 31 +++ src/backends/ccl/cncl/impl/comm_init_rank.h | 220 +++++++++++++++++++- src/backends/ccl/cncl/impl/recv.h | 17 ++ src/backends/ccl/cncl/impl/send.h | 17 ++ src/backends/mpi/ompi/impl/comm_init_all.h | 1 + src/base/comm_init_all.h | 10 +- src/base/comm_init_rank.h | 5 +- src/communicator.h | 9 + 9 files changed, 375 insertions(+), 15 deletions(-) create mode 100644 src/backends/ccl/cncl/checks.h create mode 100644 src/backends/ccl/cncl/impl/recv.h create mode 100644 src/backends/ccl/cncl/impl/send.h diff --git a/src/backends/ccl/cncl/api.h b/src/backends/ccl/cncl/api.h index b2b8ca2..a8c5469 100644 --- a/src/backends/ccl/cncl/api.h +++ b/src/backends/ccl/cncl/api.h @@ -6,6 +6,7 @@ #include #include "backends/ccl/common/api.h" +#include "devices/cambricon/checks.h" #include "logging.h" #include "return_status_impl.h" #include "runtime.h" @@ -24,6 +25,59 @@ struct CnclApi { using RedOp = cnclReduceOp_t; using Stream = typename Runtime::Stream; + private: + // CNCL does not support a `nullptr` queue argument, so we need to manage a default queue for synchronous operations. + struct DefaultQueue { + Stream queue = nullptr; + int device_id = -1; + + ~DefaultQueue() { + if (queue != nullptr) { + INFINI_CHECK_CNRT(cnrtQueueDestroy(queue)); + } + } + + void SetDevice(int new_device_id) { + if (queue != nullptr) { + INFINI_CHECK_CNRT(cnrtQueueDestroy(queue)); + } + INFINI_CHECK_CNRT(cnrtQueueCreate(&queue)); + device_id = new_device_id; + } + }; + + static Stream GetDefaultQueue() { + thread_local DefaultQueue default_queue; + + int device_id = 0; + INFINI_CHECK_CNRT(cnrtGetDevice(&device_id)); + if (default_queue.device_id != device_id) { + default_queue.SetDevice(device_id); + } + return default_queue.queue; + } + + using PointToPointOp = Result (*)(void*, size_t, DataType, int, Comm, Stream); + + static Result PointToPoint(PointToPointOp operation, void* buffer, + size_t count, DataType data_type, int peer, + Comm comm, Stream stream) { + // Up to CNCL 1.30.8, a `nullptr` queue argument is unsupported. + // Synchronize the fallback queue to preserve the synchronous behavior of the `nullptr` path; + // explicit queues stay async. Future CNCL versions may remove this compatibility path. + const bool is_default_queue = stream == nullptr; + if (is_default_queue) { + stream = GetDefaultQueue(); + } + + Result result = operation(buffer, count, data_type, peer, comm, stream); + if (is_default_queue) { + INFINI_CHECK_CNRT(cnrtQueueSync(stream)); + } + return result; + } + + public: static ReturnStatus Check(Result result) { if (result != CNCL_RET_SUCCESS) { LOG(cnclGetErrorStr(result)); @@ -34,15 +88,19 @@ struct CnclApi { static Result GetUniqueId(UniqueId* id) { return cnclGetCliqueId(id); } + static Result InitComms(Comm* comms, int num_comm, const int* dev_list, + const int* rank_list, int nrank, + UniqueId* clique_id) { + return cnclInitComms(comms, num_comm, dev_list, rank_list, nrank, + clique_id); + } + static Result CommInitRank(Comm* comm, int nranks, UniqueId id, int rank) { using Rt = Runtime; int device_id = 0; - if (Rt::GetDevice(&device_id) != cnrtSuccess) { - return CNCL_RET_ERR_MLU_RUNTIME; - } - - return cnclInitComms(comm, 1, &device_id, &rank, nranks, &id); + INFINI_CHECK_CNRT(Rt::GetDevice(&device_id)); + return InitComms(comm, 1, &device_id, &rank, nranks, &id); } static Result CommDestroy(Comm comm) { return cnclFreeComm(comm); } @@ -60,6 +118,18 @@ struct CnclApi { return cnclAllGather(send_buff, recv_buff, send_count, data_type, comm, stream); } + + static Result Send(const void* send_buff, size_t count, DataType data_type, + int peer, Comm comm, Stream stream) { + return PointToPoint(cnclSend, const_cast(send_buff), count, + data_type, peer, comm, stream); + } + + static Result Recv(void* recv_buff, size_t count, DataType data_type, + int peer, Comm comm, Stream stream) { + return PointToPoint(cnclRecv, recv_buff, count, data_type, peer, comm, + stream); + } }; } // namespace infini::ccl diff --git a/src/backends/ccl/cncl/checks.h b/src/backends/ccl/cncl/checks.h new file mode 100644 index 0000000..4e853fe --- /dev/null +++ b/src/backends/ccl/cncl/checks.h @@ -0,0 +1,31 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_CHECKS_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_CHECKS_H_ + +#include + +#include + +#include "return_status_impl.h" + +#define INFINI_CHECK_CNCL(result) \ + ::infini::ccl::detail::CheckCnclImpl((result), __FILE__, __LINE__) + +namespace infini::ccl { + +namespace detail { + +inline ReturnStatus CheckCnclImpl(cnclResult_t cncl_result, const char *file, + int line) { + if (cncl_result != CNCL_RET_SUCCESS) { + std::cerr << "backend(cncl) CNCL error code: " << cncl_result << " at line " + << line << " in " << file << std::endl; + std::abort(); + } + return ReturnStatus::kSuccess; +} + +} // namespace detail + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_CHECKS_H_ diff --git a/src/backends/ccl/cncl/impl/comm_init_rank.h b/src/backends/ccl/cncl/impl/comm_init_rank.h index 8ca3f59..49fc2c8 100644 --- a/src/backends/ccl/cncl/impl/comm_init_rank.h +++ b/src/backends/ccl/cncl/impl/comm_init_rank.h @@ -1,13 +1,227 @@ #ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ #define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ -#include "backends/ccl/common/impl/comm_init_rank.h" +#include +#include +#include +#include +#include +#include + +#include "backends/ccl/cncl/checks.h" +#include "backends/ccl/common/comm_instance.h" +#include "base/comm_init_rank.h" +#include "communicator.h" namespace infini::ccl { template -class CommInitRankImpl - : public CclCommInitRankImpl {}; +class CommInitRankImpl { + using Api = CclApi; + using CommInstance = CclCommInstance; + using CliqueId = typename Api::UniqueId; + + // CNCL initializes all communicators belonging to one process in a single + // call, so each thread contributes one request to a local initialization + // group. + struct Request { + Communicator* comm; + CliqueId clique_id; + int nranks; + int rank; + int expected_local_ranks; + ReturnStatus status = ReturnStatus::kSuccess; + bool done = false; + }; + + // Only one clique group may be initialized at a time in a process. Requests + // for another group wait in `deferred_requests` until the active group ends. + struct ThreadCoordinator { + std::mutex mutex; + std::condition_variable condition; + std::vector active_requests; + std::vector deferred_requests; + CliqueId active_clique_id{}; + int active_nranks = 0; + int active_expected_local_ranks = 0; + bool has_active_group = false; + }; + + static bool SameClique(const CliqueId& left, const CliqueId& right) { + return std::memcmp(&left, &right, sizeof(CliqueId)) == 0; + } + + static ReturnStatus ValidateRequests(const std::vector& requests, + int nranks, + int expected_local_ranks) { + // The arrays passed to `cnclInitComms` describe only this process's local + // communicators, while `nranks` describes the global communicator. + if (requests.size() != static_cast(expected_local_ranks)) { + return ReturnStatus::kInvalidArgument; + } + + for (size_t i = 0; i < requests.size(); ++i) { + if (!requests[i]->comm || requests[i]->comm->intra_comm() || + requests[i]->rank < 0 || requests[i]->rank >= nranks) { + return ReturnStatus::kInvalidArgument; + } + if (i > 0 && requests[i]->rank == requests[i - 1]->rank) { + return ReturnStatus::kInvalidArgument; + } + } + + std::vector devices; + devices.reserve(requests.size()); + for (const Request* request : requests) { + devices.push_back(request->comm->device_id()); + } + std::sort(devices.begin(), devices.end()); + for (size_t i = 1; i < devices.size(); ++i) { + if (devices[i] == devices[i - 1]) { + return ReturnStatus::kInvalidArgument; + } + } + + return ReturnStatus::kSuccess; + } + + static bool IsSameGroup(const Request& request, + const ThreadCoordinator& coordinator) { + return SameClique(request.clique_id, coordinator.active_clique_id) && + request.nranks == coordinator.active_nranks && + request.expected_local_ranks == + coordinator.active_expected_local_ranks; + } + + static int ExpectedLocalRanks(const Communicator* comm, int nranks) { + if (!comm->inter_comm()) { + return nranks; + } + + if (nranks == comm->size()) { + return 1; + } + + int local_size = comm->local_size(); + return local_size > 0 ? local_size : 1; + } + + static void StartDeferredGroup(ThreadCoordinator& coordinator) { + Request* first = coordinator.deferred_requests.front(); + coordinator.active_clique_id = first->clique_id; + coordinator.active_nranks = first->nranks; + coordinator.active_expected_local_ranks = first->expected_local_ranks; + coordinator.has_active_group = true; + + std::vector remaining; + for (Request* request : coordinator.deferred_requests) { + if (IsSameGroup(*request, coordinator)) { + coordinator.active_requests.push_back(request); + } else { + remaining.push_back(request); + } + } + coordinator.deferred_requests = std::move(remaining); + } + + static void CompleteGroup(ThreadCoordinator& coordinator) { + std::vector requests = coordinator.active_requests; + // Keep the arrays deterministic and align each returned CNCL handle with + // the communicator for the corresponding global rank. + std::sort(requests.begin(), requests.end(), + [](const Request* left, const Request* right) { + return left->rank < right->rank; + }); + + ReturnStatus status = ValidateRequests( + requests, coordinator.active_nranks, + coordinator.active_expected_local_ranks); + + std::vector comms(requests.size()); + std::vector devices(requests.size()); + std::vector ranks(requests.size()); + for (size_t i = 0; i < requests.size(); ++i) { + ranks[i] = requests[i]->rank; + devices[i] = requests[i]->comm->device_id(); + } + + if (status == ReturnStatus::kSuccess) { + INFINI_CHECK_CNCL(Api::InitComms( + comms.data(), static_cast(comms.size()), devices.data(), + ranks.data(), coordinator.active_nranks, + &requests.front()->clique_id)); + } + + for (size_t i = 0; i < requests.size(); ++i) { + Request* request = requests[i]; + if (status == ReturnStatus::kSuccess) { + request->comm->set_world_info(request->rank, coordinator.active_nranks); + auto instance = std::make_unique(); + instance->handle = comms[i]; + request->comm->set_intra_comm(std::move(instance)); + } + request->status = status; + request->done = true; + } + + coordinator.active_requests.clear(); + coordinator.has_active_group = false; + + if (!coordinator.deferred_requests.empty()) { + StartDeferredGroup(coordinator); + } + coordinator.condition.notify_all(); + } + + public: + static ReturnStatus Apply(Communicator* comm, int nranks, + infinicclUniqueId id, int rank) { + if (!comm || comm->intra_comm() || nranks <= 0 || rank < 0 || + rank >= nranks) { + return ReturnStatus::kInvalidArgument; + } + + CliqueId clique_id{}; + std::memcpy(&clique_id, id.internal, sizeof(clique_id)); + + int expected_local_ranks = ExpectedLocalRanks(comm, nranks); + if (expected_local_ranks <= 0 || expected_local_ranks > nranks) { + return ReturnStatus::kInvalidArgument; + } + + // CNCL accepts each clique ID only once per process. Local ranks for the + // same process rendezvous here so one `cnclInitComms` call creates the + // process-local communicator set. + static ThreadCoordinator coordinator; + Request request{comm, clique_id, nranks, rank, expected_local_ranks}; + + std::unique_lock lock(coordinator.mutex); + if (!coordinator.has_active_group) { + coordinator.active_clique_id = clique_id; + coordinator.active_nranks = nranks; + coordinator.active_expected_local_ranks = expected_local_ranks; + coordinator.has_active_group = true; + } + + if (IsSameGroup(request, coordinator)) { + coordinator.active_requests.push_back(&request); + } else { + coordinator.deferred_requests.push_back(&request); + } + + while (!request.done) { + if (IsSameGroup(request, coordinator) && + static_cast(coordinator.active_requests.size()) == + expected_local_ranks) { + CompleteGroup(coordinator); + } else { + coordinator.condition.wait(lock); + } + } + + return request.status; + } +}; template <> struct BackendEnabled : std::true_type {}; diff --git a/src/backends/ccl/cncl/impl/recv.h b/src/backends/ccl/cncl/impl/recv.h new file mode 100644 index 0000000..0ad9cdd --- /dev/null +++ b/src/backends/ccl/cncl/impl/recv.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_RECV_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_RECV_H_ + +#include "backends/ccl/common/impl/recv.h" + +namespace infini::ccl { + +template +class RecvImpl + : public CclRecvImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_RECV_H_ diff --git a/src/backends/ccl/cncl/impl/send.h b/src/backends/ccl/cncl/impl/send.h new file mode 100644 index 0000000..0ed2698 --- /dev/null +++ b/src/backends/ccl/cncl/impl/send.h @@ -0,0 +1,17 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_SEND_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_SEND_H_ + +#include "backends/ccl/common/impl/send.h" + +namespace infini::ccl { + +template +class SendImpl + : public CclSendImpl {}; + +template <> +struct BackendEnabled : std::true_type {}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_SEND_H_ diff --git a/src/backends/mpi/ompi/impl/comm_init_all.h b/src/backends/mpi/ompi/impl/comm_init_all.h index 3217655..4812f9c 100644 --- a/src/backends/mpi/ompi/impl/comm_init_all.h +++ b/src/backends/mpi/ompi/impl/comm_init_all.h @@ -31,6 +31,7 @@ class CommInitAllImpl { INFINI_CHECK_MPI(MPI_Comm_size(inst->handle, &size)); comm->set_world_info(rank, size); + comm->set_local_size(n_dev); comm->set_inter_comm(std::move(inst)); int local_rank = 0; diff --git a/src/base/comm_init_all.h b/src/base/comm_init_all.h index c37d781..c96e1c9 100644 --- a/src/base/comm_init_all.h +++ b/src/base/comm_init_all.h @@ -13,10 +13,10 @@ struct CommInitAllImpl; class CommInitAll : public Operation { public: - template - static ReturnStatus Execute(void** comm_handles, int n_dev, - const int* dev_list) { - if (!comm_handles || n_dev <= 0) { + template + static ReturnStatus Execute(void** comm_handle, int n_dev, + Args &&...args) { + if (!comm_handle || n_dev <= 0) { LOG("Invalid communicator handle for `CommInitAll`."); return ReturnStatus::kInvalidArgument; } @@ -31,7 +31,7 @@ class CommInitAll : public Operation { } return CommInitAllImpl::Apply( - comm, std::forward(args)...); + comm, n_dev, std::forward(args)...); } }; diff --git a/src/base/comm_init_rank.h b/src/base/comm_init_rank.h index e4391d4..d420a6b 100644 --- a/src/base/comm_init_rank.h +++ b/src/base/comm_init_rank.h @@ -1,6 +1,7 @@ #ifndef INFINI_CCL_BASE_COMM_INIT_RANK_H_ #define INFINI_CCL_BASE_COMM_INIT_RANK_H_ +#include "communicator.h" #include "logging.h" #include "operation.h" #include "return_status_impl.h" @@ -14,9 +15,9 @@ class CommInitRank : public Operation { public: template - static ReturnStatus Execute(void **comm_handle, Args &&...args) { + static ReturnStatus Execute(void **comm_handle, Args &&...args) { if (!comm_handle) { - LOG("Invalid communicator handle for `CommInitAll`."); + LOG("Invalid communicator handle for `CommInitRank`."); return ReturnStatus::kInvalidArgument; } diff --git a/src/communicator.h b/src/communicator.h index fe70c87..0c4fee3 100644 --- a/src/communicator.h +++ b/src/communicator.h @@ -53,6 +53,13 @@ class Communicator { int size() const { return global_size_; } + int local_size() const { return local_size_; } + + void set_local_size(int size) { + local_size_ = size; + return; + } + int device_id() const { return device_id_; } void set_device_id(int id) { @@ -79,6 +86,8 @@ class Communicator { int global_size_; + int local_size_ = 0; + Device::Type device_type_; }; From c2cb39bfcfedf00f7bb73e709687c80c497973ff Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Wed, 23 Sep 2026 11:15:13 +0000 Subject: [PATCH 4/7] style: apply clang-format to the unformatted files --- src/backends/ccl/cncl/api.h | 10 ++++++---- src/backends/ccl/cncl/impl/comm_init_rank.h | 11 +++++------ src/base/comm_init_all.h | 6 +++--- 3 files changed, 14 insertions(+), 13 deletions(-) diff --git a/src/backends/ccl/cncl/api.h b/src/backends/ccl/cncl/api.h index a8c5469..173f0eb 100644 --- a/src/backends/ccl/cncl/api.h +++ b/src/backends/ccl/cncl/api.h @@ -26,7 +26,8 @@ struct CnclApi { using Stream = typename Runtime::Stream; private: - // CNCL does not support a `nullptr` queue argument, so we need to manage a default queue for synchronous operations. + // CNCL does not support a `nullptr` queue argument, so we need to manage a + // default queue for synchronous operations. struct DefaultQueue { Stream queue = nullptr; int device_id = -1; @@ -62,9 +63,10 @@ struct CnclApi { static Result PointToPoint(PointToPointOp operation, void* buffer, size_t count, DataType data_type, int peer, Comm comm, Stream stream) { - // Up to CNCL 1.30.8, a `nullptr` queue argument is unsupported. - // Synchronize the fallback queue to preserve the synchronous behavior of the `nullptr` path; - // explicit queues stay async. Future CNCL versions may remove this compatibility path. + // Up to CNCL 1.30.8, a `nullptr` queue argument is unsupported. + // Synchronize the fallback queue to preserve the synchronous behavior of + // the `nullptr` path; explicit queues stay async. Future CNCL versions may + // remove this compatibility path. const bool is_default_queue = stream == nullptr; if (is_default_queue) { stream = GetDefaultQueue(); diff --git a/src/backends/ccl/cncl/impl/comm_init_rank.h b/src/backends/ccl/cncl/impl/comm_init_rank.h index 49fc2c8..ed33fc8 100644 --- a/src/backends/ccl/cncl/impl/comm_init_rank.h +++ b/src/backends/ccl/cncl/impl/comm_init_rank.h @@ -52,8 +52,7 @@ class CommInitRankImpl { } static ReturnStatus ValidateRequests(const std::vector& requests, - int nranks, - int expected_local_ranks) { + int nranks, int expected_local_ranks) { // The arrays passed to `cnclInitComms` describe only this process's local // communicators, while `nranks` describes the global communicator. if (requests.size() != static_cast(expected_local_ranks)) { @@ -133,10 +132,10 @@ class CommInitRankImpl { return left->rank < right->rank; }); - ReturnStatus status = ValidateRequests( - requests, coordinator.active_nranks, - coordinator.active_expected_local_ranks); - + ReturnStatus status = + ValidateRequests(requests, coordinator.active_nranks, + coordinator.active_expected_local_ranks); + std::vector comms(requests.size()); std::vector devices(requests.size()); std::vector ranks(requests.size()); diff --git a/src/base/comm_init_all.h b/src/base/comm_init_all.h index c96e1c9..df24941 100644 --- a/src/base/comm_init_all.h +++ b/src/base/comm_init_all.h @@ -13,9 +13,9 @@ struct CommInitAllImpl; class CommInitAll : public Operation { public: - template - static ReturnStatus Execute(void** comm_handle, int n_dev, - Args &&...args) { + template + static ReturnStatus Execute(void **comm_handle, int n_dev, Args &&...args) { if (!comm_handle || n_dev <= 0) { LOG("Invalid communicator handle for `CommInitAll`."); return ReturnStatus::kInvalidArgument; From ccac36d7c26c0a8aa22227405b82988d2588847c Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Wed, 23 Sep 2026 20:07:48 +0800 Subject: [PATCH 5/7] refactor: use `ConstexprMap` for CNCL's type maps --- src/backends/ccl/cncl/type_map.h | 75 ++++++++++++++++---------------- 1 file changed, 37 insertions(+), 38 deletions(-) diff --git a/src/backends/ccl/cncl/type_map.h b/src/backends/ccl/cncl/type_map.h index 04f6eee..bff1330 100644 --- a/src/backends/ccl/cncl/type_map.h +++ b/src/backends/ccl/cncl/type_map.h @@ -3,50 +3,44 @@ #include +#include + #include "backends/ccl/common/api.h" #include "comm_impl.h" #include "data_type_impl.h" +#include "logging.h" namespace infini::ccl { -inline bool DataTypeToCnclType(DataType dtype, cnclDataType_t* cncl_dtype) { - switch (dtype) { - case DataType::kInt8: - *cncl_dtype = cnclInt8; - return true; - case DataType::kInt16: - *cncl_dtype = cnclInt16; - return true; - case DataType::kInt32: - *cncl_dtype = cnclInt32; - return true; - case DataType::kInt64: - *cncl_dtype = cnclInt64; - return true; - case DataType::kUInt8: - *cncl_dtype = cnclUint8; - return true; - case DataType::kUInt16: - *cncl_dtype = cnclUint16; - return true; - case DataType::kUInt32: - *cncl_dtype = cnclUint32; - return true; - case DataType::kUInt64: - *cncl_dtype = cnclUint64; - return true; - case DataType::kFloat16: - *cncl_dtype = cnclFloat16; - return true; - case DataType::kBFloat16: - *cncl_dtype = cnclBfloat16; - return true; - case DataType::kFloat32: - *cncl_dtype = cnclFloat32; - return true; - default: - return false; +template +struct CnclDataTypeMap { + static constexpr ConstexprMap kMap{{{ + {DataType::kInt8, cnclInt8}, + {DataType::kInt16, cnclInt16}, + {DataType::kInt32, cnclInt32}, + {DataType::kInt64, cnclInt64}, + {DataType::kUInt8, cnclUint8}, + {DataType::kUInt16, cnclUint16}, + {DataType::kUInt32, cnclUint32}, + {DataType::kUInt64, cnclUint64}, + {DataType::kFloat32, cnclFloat32}, + {DataType::kFloat64, cnclInvalid}, + {DataType::kFloat16, cnclFloat16}, + {DataType::kBFloat16, cnclBfloat16}, + }}}; +}; + +template +inline cnclDataType_t DataTypeToCnclType(DataType dtype) { + auto cncl_dtype = CnclDataTypeMap::kMap.at(dtype); + + if (cncl_dtype == cnclInvalid) { + LOG(("DataType '" + std::string(kDataTypeToDesc.at(dtype)) + + "' is not supported by the CNCL backend") + .c_str()); } + + return cncl_dtype; } template <> @@ -55,7 +49,12 @@ struct CclTypeMap { static bool ToBackendDataType(DataType dtype, typename Api::DataType* backend_dtype) { - return DataTypeToCnclType(dtype, backend_dtype); + auto cncl_dtype = DataTypeToCnclType(dtype); + if (cncl_dtype == cnclInvalid) { + return false; + } + *backend_dtype = cncl_dtype; + return true; } static bool ToBackendRedOp(ReductionOpType red_op, From ebad8952e283c7e901326f3e6b7c769e54b7bad7 Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Wed, 23 Sep 2026 20:24:42 +0800 Subject: [PATCH 6/7] refactor: add the CNCL type map for `RedOp` --- src/backends/ccl/cncl/type_map.h | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/src/backends/ccl/cncl/type_map.h b/src/backends/ccl/cncl/type_map.h index bff1330..bca28a6 100644 --- a/src/backends/ccl/cncl/type_map.h +++ b/src/backends/ccl/cncl/type_map.h @@ -43,6 +43,17 @@ inline cnclDataType_t DataTypeToCnclType(DataType dtype) { return cncl_dtype; } +static const ConstexprMap kCnclOpMap{{{ + {ReductionOpType::kSum, cnclSum}, + {ReductionOpType::kProd, cnclProd}, + {ReductionOpType::kMax, cnclMax}, + {ReductionOpType::kMin, cnclMin}, +}}}; + +inline cnclReduceOp_t RedOpToCnclOp(ReductionOpType red_op) { + return kCnclOpMap.at(red_op); +} + template <> struct CclTypeMap { using Api = CclApi; @@ -62,7 +73,7 @@ struct CclTypeMap { if (red_op == ReductionOpType::kAvg) { return false; } - *backend_op = static_cast(red_op); + *backend_op = RedOpToCnclOp(red_op); return true; } }; From e7de95c57d32614300e19fe0a4fb3a6d3bf84538 Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Wed, 23 Sep 2026 20:50:11 +0800 Subject: [PATCH 7/7] docs: update `README.md` and the PR template to include CNCL as a supported backend --- .github/pull_request_template.md | 2 ++ README.md | 2 ++ 2 files changed, 4 insertions(+) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index cfe73dd..326ede4 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -66,6 +66,7 @@ Please check all the platforms and/or backends this PR affects (i.e., code is to - [ ] MPICH - [ ] NCCL/RCCL - [ ] MCCL +- [ ] CNCL ## Performance Impact @@ -116,6 +117,7 @@ See `CONTRIBUTING.md` ยง Pull Requests for the official testing requirements and - [ ] MPICH - [ ] NCCL/RCCL - [ ] MCCL +- [ ] CNCL --- diff --git a/README.md b/README.md index 4a694d7..0158be0 100644 --- a/README.md +++ b/README.md @@ -144,6 +144,7 @@ cmake .. -DWITH_NVIDIA=ON -DWITH_OMPI=ON | `WITH_MPICH` | Enable MPICH backend | `OFF` | | `WITH_NCCL` | Enable NCCL/RCCL backend | `OFF` | | `WITH_MCCL` | Enable MCCL backend | `OFF` | +| `WITH_CNCL` | Enable CNCL backend | `OFF` | | **Miscellaneous** ||| | `AUTO_DETECT_DEVICES` | Automatically detect available devices and enable corresponding support | `ON` | | `AUTO_DETECT_BACKENDS` | Automatically detect available communication backends and enable corresponding support | `OFF` | @@ -357,6 +358,7 @@ export LD_LIBRARY_PATH=${INFINI_INSTALL}/lib:$LD_LIBRARY_PATH | **MPICH** | Full | `WITH_MPICH=ON` | Requires the MPICH development package.| | **NCCL** | Partial | `WITH_NCCL=ON` | Requires NVIDIA or Iluvatar NCCL, or HYGON RCCL. Currently available when `WITH_NVIDIA=ON`, `WITH_ILUVATAR=ON`, or `WITH_HYGON=ON`.| | **MCCL** | Partial | `WITH_MCCL=ON` | Requires MetaX or Moore MCCL. Currently available when `WITH_METAX=ON` or `WITH_MOORE=ON`.| +| **CNCL** | Partial | `WITH_CNCL=ON` | Requires Cambricon CNCL. Available only when `WITH_CAMBRICON=ON`.|