From b4eedddaead234801c03ad82fda66d41f5ee4cc8 Mon Sep 17 00:00:00 2001 From: Zimin Li Date: Tue, 4 Aug 2026 15:52:29 +0800 Subject: [PATCH] fix: gate Moore MCCL bfloat16 support by `MARCH_TYPE` --- examples/CMakeLists.txt | 3 --- src/CMakeLists.txt | 4 --- src/backends/ccl/mccl/api.h | 3 +++ src/backends/ccl/mccl/metax/api.h | 5 ++++ src/backends/ccl/mccl/moore/api.h | 9 +++++++ src/backends/ccl/mccl/type_map.h | 43 +++++++++++++++---------------- 6 files changed, 38 insertions(+), 29 deletions(-) diff --git a/examples/CMakeLists.txt b/examples/CMakeLists.txt index 2a31f97..24495ea 100644 --- a/examples/CMakeLists.txt +++ b/examples/CMakeLists.txt @@ -43,9 +43,6 @@ foreach(source_file ${EXAMPLE_SOURCES}) if(WITH_MOORE) target_link_libraries(${target_name} PRIVATE ${MUSART_LIB}) target_compile_options(${target_name} PRIVATE "-x" "musa") - if(WITH_MCCL) - target_compile_definitions(${target_name} PRIVATE INFINI_CCL_MCCL_BFLOAT16_UNSUPPORTED) - endif() endif() if(WITH_CAMBRICON) diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 9584156..06df487 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -243,10 +243,6 @@ if(WITH_MCCL) target_sources(infiniccl PRIVATE ${MCCL_SRCS}) target_include_directories(infiniccl PRIVATE ${MCCL_INC}) target_link_libraries(infiniccl PRIVATE ${MCCL_LIB}) - - if(WITH_MOORE) - target_compile_definitions(infiniccl PRIVATE INFINI_CCL_MCCL_BFLOAT16_UNSUPPORTED) - endif() endif() # ========================================================= diff --git a/src/backends/ccl/mccl/api.h b/src/backends/ccl/mccl/api.h index 06398d2..a5d2bcd 100644 --- a/src/backends/ccl/mccl/api.h +++ b/src/backends/ccl/mccl/api.h @@ -12,6 +12,9 @@ namespace infini::ccl { +template +struct McclDataTypeTraits; + template struct McclApi { static constexpr BackendType kBackendType = BackendType::kMccl; diff --git a/src/backends/ccl/mccl/metax/api.h b/src/backends/ccl/mccl/metax/api.h index 0817a01..fba024e 100644 --- a/src/backends/ccl/mccl/metax/api.h +++ b/src/backends/ccl/mccl/metax/api.h @@ -6,6 +6,11 @@ namespace infini::ccl { +template <> +struct McclDataTypeTraits { + static constexpr mcclDataType_t kBFloat16 = mcclBfloat16; +}; + template <> struct CclApi : McclApi {}; diff --git a/src/backends/ccl/mccl/moore/api.h b/src/backends/ccl/mccl/moore/api.h index 198f961..e4860d1 100644 --- a/src/backends/ccl/mccl/moore/api.h +++ b/src/backends/ccl/mccl/moore/api.h @@ -6,6 +6,15 @@ namespace infini::ccl { +template <> +struct McclDataTypeTraits { +#if defined(MARCH_TYPE) && MARCH_TYPE >= 220 + static constexpr mcclDataType_t kBFloat16 = mcclBfloat16; +#else + static constexpr mcclDataType_t kBFloat16 = mcclNumTypes; +#endif +}; + template <> struct CclApi : McclApi {}; diff --git a/src/backends/ccl/mccl/type_map.h b/src/backends/ccl/mccl/type_map.h index 250445b..4364ba0 100644 --- a/src/backends/ccl/mccl/type_map.h +++ b/src/backends/ccl/mccl/type_map.h @@ -6,32 +6,30 @@ #include #include "backends/ccl/common/api.h" +#include "backends/ccl/mccl/api.h" #include "comm_impl.h" #include "data_type_impl.h" #include "logging.h" namespace infini::ccl { -#if defined(INFINI_CCL_MCCL_BFLOAT16_UNSUPPORTED) -constexpr mcclDataType_t kMcclBFloat16Val = mcclNumTypes; -#else -constexpr mcclDataType_t kMcclBFloat16Val = mcclBfloat16; -#endif - -static const ConstexprMap kMcclTypeMap{{{ - {DataType::kInt8, mcclInt8}, - {DataType::kInt16, mcclNumTypes}, - {DataType::kInt32, mcclInt32}, - {DataType::kInt64, mcclInt64}, - {DataType::kUInt8, mcclUint8}, - {DataType::kUInt16, mcclNumTypes}, - {DataType::kUInt32, mcclUint32}, - {DataType::kUInt64, mcclUint64}, - {DataType::kFloat32, mcclFloat32}, - {DataType::kFloat64, mcclFloat64}, - {DataType::kFloat16, mcclFloat16}, - {DataType::kBFloat16, kMcclBFloat16Val}, -}}}; +template +struct McclDataTypeMap { + static constexpr ConstexprMap kMap{{{ + {DataType::kInt8, mcclInt8}, + {DataType::kInt16, mcclNumTypes}, + {DataType::kInt32, mcclInt32}, + {DataType::kInt64, mcclInt64}, + {DataType::kUInt8, mcclUint8}, + {DataType::kUInt16, mcclNumTypes}, + {DataType::kUInt32, mcclUint32}, + {DataType::kUInt64, mcclUint64}, + {DataType::kFloat32, mcclFloat32}, + {DataType::kFloat64, mcclFloat64}, + {DataType::kFloat16, mcclFloat16}, + {DataType::kBFloat16, McclDataTypeTraits::kBFloat16}, + }}}; +}; static const ConstexprMap kMcclOpMap{{{ {ReductionOpType::kSum, mcclSum}, @@ -41,8 +39,9 @@ static const ConstexprMap kMcclOpMap{{{ {ReductionOpType::kAvg, mcclAvg}, }}}; +template inline mcclDataType_t DataTypeToMcclType(DataType dtype) { - auto mccl_dtype = kMcclTypeMap.at(dtype); + auto mccl_dtype = McclDataTypeMap::kMap.at(dtype); if (mccl_dtype == mcclNumTypes) { LOG(("DataType '" + std::string(kDataTypeToDesc.at(dtype)) + @@ -63,7 +62,7 @@ struct CclTypeMap { static bool ToBackendDataType(DataType dtype, typename Api::DataType *backend_dtype) { - auto mccl_dtype = DataTypeToMcclType(dtype); + auto mccl_dtype = DataTypeToMcclType(dtype); if (mccl_dtype == mcclNumTypes) { return false; }