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/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/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`.| 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/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..173f0eb --- /dev/null +++ b/src/backends/ccl/cncl/api.h @@ -0,0 +1,139 @@ +#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 "devices/cambricon/checks.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 UniqueId = cnclCliqueId; + using Result = cnclResult_t; + using DataType = cnclDataType_t; + 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)); + return ReturnStatus::kSystemError; + } + return ReturnStatus::kSuccess; + } + + 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; + INFINI_CHECK_CNRT(Rt::GetDevice(&device_id)); + return InitComms(comm, 1, &device_id, &rank, nranks, &id); + } + + 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); + } + + 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 + +#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/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/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_rank.h b/src/backends/ccl/cncl/impl/comm_init_rank.h new file mode 100644 index 0000000..ed33fc8 --- /dev/null +++ b/src/backends/ccl/cncl/impl/comm_init_rank.h @@ -0,0 +1,230 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_IMPL_COMM_INIT_RANK_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_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 { + 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 {}; + +} // 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/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/ccl/cncl/type_map.h b/src/backends/ccl/cncl/type_map.h new file mode 100644 index 0000000..bca28a6 --- /dev/null +++ b/src/backends/ccl/cncl/type_map.h @@ -0,0 +1,83 @@ +#ifndef INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_H_ +#define INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_H_ + +#include + +#include + +#include "backends/ccl/common/api.h" +#include "comm_impl.h" +#include "data_type_impl.h" +#include "logging.h" + +namespace infini::ccl { + +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; +} + +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; + + static bool ToBackendDataType(DataType dtype, + typename Api::DataType* 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, + typename Api::RedOp* backend_op) { + if (red_op == ReductionOpType::kAvg) { + return false; + } + *backend_op = RedOpToCnclOp(red_op); + return true; + } +}; + +} // namespace infini::ccl + +#endif // INFINI_CCL_BACKENDS_CCL_CNCL_TYPE_MAP_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 e834622..4812f9c 100644 --- a/src/backends/mpi/ompi/impl/comm_init_all.h +++ b/src/backends/mpi/ompi/impl/comm_init_all.h @@ -18,11 +18,10 @@ class CommInitAllImpl { ListGetBest(ActiveDevices{}); using Rt = Runtime; - if (!comm) { + if (comm && comm->inter_comm()) { // TODO(lzm): change to use `glog`. - LOG("Failed to initialize OpenMPI communicator: invalid " - "communicator pointer."); - return ReturnStatus::kInternalError; + LOG("Invalid communicator handle for `CommInitAll`."); + return ReturnStatus::kInvalidArgument; } int rank, size; @@ -32,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 e737306..df24941 100644 --- a/src/base/comm_init_all.h +++ b/src/base/comm_init_all.h @@ -15,10 +15,8 @@ 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`. + 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; } @@ -26,12 +24,14 @@ class CommInitAll : public Operation { 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)...); + comm, n_dev, std::forward(args)...); } }; diff --git a/src/base/comm_init_rank.h b/src/base/comm_init_rank.h index 46f93f2..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" @@ -15,9 +16,7 @@ class CommInitRank : public Operation { 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`. + if (!comm_handle) { LOG("Invalid communicator handle for `CommInitRank`."); return ReturnStatus::kInvalidArgument; } @@ -29,6 +28,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 { 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_; };