LCOV - code coverage report
Current view: top level - legacy/ascend950/framework/communicator - op_params_checker.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 87.8 % 147 129
Test Date: 2026-08-04 10:52:23 Functions: 91.7 % 12 11

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "op_params_checker.h"
      12              : #include <string>
      13              : #include "hccl_params_pub.h"
      14              : #include "data_type.h"
      15              : #include "op_type.h"
      16              : #include "string_util.h"
      17              : #include "exception_util.h"
      18              : #include "adapter_error_manager_pub.h"
      19              : 
      20              : namespace Hccl {
      21              : 
      22          334 : HcclResult OpParamsChecker::CheckOpDataTypeOpbase(const CollOpParams &opParams, bool ccuEnable, bool isDevUsed, bool isAiv)
      23              : {
      24          334 :     HcclResult ret = HcclResult::HCCL_E_PARA;
      25          334 :     if (ccuEnable){
      26          143 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapCcuOpbase);
      27          191 :     } else if (isDevUsed) {
      28          190 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapAicpuOpbase);
      29            1 :     } else if (isAiv) {
      30            0 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapAivOpbase);
      31              :     } else {
      32            3 :         HCCL_ERROR("[OpParamsChecker::%s] Host opbase mode is invalid.", __func__);
      33              :     }
      34          334 :     return ret;
      35              : }
      36              : 
      37          320 : HcclResult OpParamsChecker::CheckOpDataTypeOffload(const CollOpParams &opParams, bool ccuEnable, bool isDevUsed, bool isAiv)
      38              : {
      39          320 :     HcclResult ret = HcclResult::HCCL_E_PARA;
      40          320 :     if (ccuEnable){
      41          104 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapCcuOffload);
      42          216 :     } else if (isDevUsed) {
      43          105 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapAicpuOffload);
      44          111 :     } else if (isAiv) {
      45           18 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapAivOffload);
      46              :     } else {
      47           93 :         ret = CheckOpDataTypeByMap(opParams, opDataTypeSupportMapHostOffload);
      48              :     }
      49          320 :     return ret;
      50              : }
      51              : 
      52            0 : static void ReportOpTypeErrMsg(const std::string& callName, OpType opType)
      53              : {
      54            0 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
      55              :             std::vector<std::string>({callName, opType.Describe(), "opType",
      56              :             "please check opType that is not supported"}));
      57            0 : }
      58              : 
      59            2 : static void ReportInputDataTypeMC2HighPErrMsg(const std::string& callName, OpType opType, DataType inputDataType)
      60              : {
      61            2 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
      62              :                     std::vector<std::string>({callName, "[" + opType.Describe() + "][" + inputDataType.Describe() + "]", "[opType][dataType]",
      63              :                     "FP32,FP16,BF16,UINT8,INT16,INT32"}));
      64            2 : }
      65              : 
      66            2 : static void ReportInputDataTypeMC2LowPErrMsg(const std::string& callName, OpType opType, DataType inputDataType)
      67              : {
      68            2 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
      69              :                     std::vector<std::string>({callName, "[" + opType.Describe() + "][" + inputDataType.Describe() + "]",
      70              :                         "[opType][inputDataType]", "Mc2LowP input:HIF8,E4M3,E5M2,INT8"}));
      71            2 : }
      72              : 
      73            2 : static void ReportOutputDataTypeMC2LowPErrMsg(const std::string& callName, OpType opType, DataType outputDataType)
      74              : {
      75            2 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
      76              :                     std::vector<std::string>({callName, "[" + opType.Describe() + "][" + outputDataType.Describe() + "]",
      77              :                          "[opType][outputDataType]", "Mc2LowP output:FP32,FP16,BF16"}));
      78            2 : }
      79              : 
      80            4 : static void ReportDataTypeNotTheSameErrMsg(const std::string& callName, OpType opType, DataType inputDataType, DataType outputDataType)
      81              : {
      82            4 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
      83              :                     std::vector<std::string>({callName, 
      84              :                         "[" + opType.Describe() + "][" + inputDataType.Describe() + "and" + outputDataType.Describe() + "]",
      85              :                         "[opType][inputDataType and outputDataType]", "should be same"}));
      86            4 : }
      87              : 
      88           61 : HcclResult OpParamsChecker::CheckOpDataTypeMC2(const Mc2CommConfig &config)
      89              : {
      90           61 :     OpType opType           = MC2OpType(static_cast<AicpuComType>(config.opType));
      91           61 :     DataType inputDataType  = MC2DataType(static_cast<HcclDataType>(config.dataType));
      92           61 :     DataType outputDataType = MC2DataType(static_cast<HcclDataType>(config.outputDataType));
      93              : 
      94              :     // 支持算子情况检验
      95           61 :     auto iter = opDataTypeSupportMapMC2.find(opType);
      96           61 :     if (iter == opDataTypeSupportMapMC2.end()) {
      97            0 :         ReportOpTypeErrMsg(__func__, opType);
      98              :         std::string msg = StringFormat("[OpParamsChecker::%s] unsupported opType [%s].",
      99            0 :                           __func__, opType.Describe().c_str());
     100            0 :         THROW<InvalidParamsException>(msg);
     101            0 :     }
     102              : 
     103              :     /* CCU数据类型校验规则
     104              :      * Reduce算子:
     105              :      *      高精度模式,当inputDataType==outputDataType时,可选类型为FP32、FP16、BF16、INT16、INT32,暂不支持UINT8;
     106              :      *      低精度模式,当inputDataType!=outputDataType时,inputDataType可选范围HIF8、E4M3、E5M2、INT8;outputDataType可选范围FP32、FP16、BF16;
     107              :      * 非Reduce算子:任意数据类型,inputDataType==outputDataType即可。
     108              :      */
     109           61 :     bool checkResult = false;
     110           61 :     if (opType == OpType::REDUCESCATTER || opType == OpType::ALLREDUCE){
     111           37 :         if (inputDataType == outputDataType){
     112           11 :             checkResult = dataTypeMC2HighP.test(static_cast<int>(inputDataType));
     113           11 :             if (!checkResult){
     114            1 :                 ReportInputDataTypeMC2HighPErrMsg(__func__, opType, inputDataType);
     115              :                 std::string msg = StringFormat("[OpParamsChecker::%s] opType [%s] not support data type [%s].",
     116            1 :                                   __func__, opType.Describe().c_str(), inputDataType.Describe().c_str());
     117            1 :                 THROW<InvalidParamsException>(msg);
     118            1 :             }
     119              :         } else {
     120           26 :             checkResult = inputDataTypeMC2LowP.test(static_cast<int>(inputDataType));
     121           26 :             if (!checkResult){
     122            1 :                 ReportInputDataTypeMC2LowPErrMsg(__func__, opType, inputDataType);
     123              :                 std::string msg = StringFormat("[OpParamsChecker::%s] Mc2LowP InputDataType[%s] != OutputDataType[%s] for OpType[%s], not support input data type [%s].",
     124            1 :                                   __func__, inputDataType.Describe().c_str(), outputDataType.Describe().c_str(), opType.Describe().c_str(), inputDataType.Describe().c_str());
     125            1 :                 THROW<InvalidParamsException>(msg);
     126            1 :             }
     127           25 :             checkResult = OutputDataTypeMC2LowP.test(static_cast<int>(outputDataType));
     128           25 :             if (!checkResult){
     129            1 :                 ReportOutputDataTypeMC2LowPErrMsg(__func__, opType, outputDataType);
     130              :                 std::string msg = StringFormat("[OpParamsChecker::%s] Mc2LowP InputDataType[%s] != OutputDataType[%s] for OpType[%s], not support output data type [%s].",
     131            1 :                                   __func__, inputDataType.Describe().c_str(), outputDataType.Describe().c_str(), opType.Describe().c_str(), outputDataType.Describe().c_str());
     132            1 :                 THROW<InvalidParamsException>(msg);
     133            1 :             }
     134              :         }
     135              :     } else {
     136           24 :         if (inputDataType != outputDataType) {
     137            4 :             ReportDataTypeNotTheSameErrMsg(__func__, opType, inputDataType, outputDataType);
     138              :             std::string msg = StringFormat("[OpParamsChecker::%s] DataType[%s] != OutputDataType[%s] for OpType[%s].",
     139            8 :                               __func__, inputDataType.Describe().c_str(),
     140           12 :                               outputDataType.Describe().c_str(), opType.Describe().c_str());
     141            4 :             THROW<InvalidParamsException>(msg);
     142            4 :         }
     143              :     }
     144           54 :     return HcclResult::HCCL_SUCCESS;
     145              : }
     146              : 
     147           58 : HcclResult OpParamsChecker::CheckOpDataTypeMC2V2(const Mc2CcTilingInner &config)
     148              : {
     149           58 :     OpType opType           = MC2OpType(static_cast<AicpuComType>(config.opType));
     150           58 :     DataType inputDataType  = MC2DataType(static_cast<HcclDataType>(config.srcDataType));
     151           58 :     DataType outputDataType = MC2DataType(static_cast<HcclDataType>(config.dstDataType));
     152              : 
     153              :     // 支持算子情况检验
     154           58 :     auto iter = opDataTypeSupportMapMC2.find(opType);
     155           58 :     if (iter == opDataTypeSupportMapMC2.end()) {
     156            0 :         ReportOpTypeErrMsg(__func__, opType);
     157              :         std::string msg = StringFormat("[OpParamsChecker::%s] unsupported opType [%s].",
     158            0 :                           __func__, opType.Describe().c_str());
     159            0 :         THROW<InvalidParamsException>(msg);
     160            0 :     }
     161              : 
     162              :     /* CCU数据类型校验规则
     163              :      * Reduce算子:
     164              :      *      高精度模式,当dataType==outputDataType时,可选类型为FP32、FP16、BF16、UINT8、INT16、INT32;
     165              :      *      低精度模式,当dataType!=outputDataType时,dataType可选范围HIF8、E4M3、E5M2、INT8;outputDataType可选范围FP32、FP16、BF16;
     166              :      * 非Reduce算子:任意数据类型,dataType==outputDataType即可。
     167              :      */
     168           58 :     bool checkResult = false;
     169           58 :     if (opType == OpType::REDUCESCATTER || opType == OpType::ALLREDUCE){
     170           37 :         if (inputDataType == outputDataType){
     171           11 :             checkResult = dataTypeMC2HighP.test(static_cast<int>(inputDataType));
     172           11 :             if (!checkResult){
     173            1 :                 ReportInputDataTypeMC2HighPErrMsg(__func__, opType, inputDataType);
     174              :                 std::string msg = StringFormat("[OpParamsChecker::%s] opType [%s] not support data type [%s].",
     175            1 :                                   __func__, opType.Describe().c_str(), inputDataType.Describe().c_str());
     176            1 :                 THROW<InvalidParamsException>(msg);
     177            1 :             }
     178              :         } else {
     179           26 :             checkResult = inputDataTypeMC2LowP.test(static_cast<int>(inputDataType));
     180           26 :             if (!checkResult){
     181            1 :                 ReportInputDataTypeMC2LowPErrMsg(__func__, opType, inputDataType);
     182              :                 std::string msg = StringFormat("[OpParamsChecker::%s] Mc2LowP InputDataType[%s] != OutputDataType[%s] for OpType[%s], not support input data type [%s].",
     183            1 :                                   __func__, inputDataType.Describe().c_str(), outputDataType.Describe().c_str(), opType.Describe().c_str(), inputDataType.Describe().c_str());
     184            1 :                 THROW<InvalidParamsException>(msg);
     185            1 :             }
     186           25 :             checkResult = OutputDataTypeMC2LowP.test(static_cast<int>(outputDataType));
     187           25 :             if (!checkResult){
     188            1 :                 ReportOutputDataTypeMC2LowPErrMsg(__func__, opType, outputDataType);
     189              :                 std::string msg = StringFormat("[OpParamsChecker::%s] Mc2LowP InputDataType[%s] != OutputDataType[%s] for OpType[%s], not support output data type [%s].",
     190            1 :                                   __func__, inputDataType.Describe().c_str(), outputDataType.Describe().c_str(), opType.Describe().c_str(), outputDataType.Describe().c_str());
     191            1 :                 THROW<InvalidParamsException>(msg);
     192            1 :             }
     193              :         }
     194              :     } else {
     195           21 :         if (inputDataType != outputDataType) {
     196            0 :             ReportDataTypeNotTheSameErrMsg(__func__, opType, inputDataType, outputDataType);
     197              :             std::string msg = StringFormat("[OpParamsChecker::%s] DataType[%s] != OutputDataType[%s] for OpType[%s].",
     198            0 :                               __func__, inputDataType.Describe().c_str(),
     199            0 :                               outputDataType.Describe().c_str(), opType.Describe().c_str());
     200            0 :             THROW<InvalidParamsException>(msg);
     201            0 :         }
     202              :     }
     203           55 :     return HcclResult::HCCL_SUCCESS;
     204              : }
     205              : 
     206          650 : DataType OpParamsChecker::GetDataType(const CollOpParams &opParams)
     207              : {
     208          650 :     DataType dtype = opParams.dataType;
     209          650 :     if (opParams.opType == OpType::ALLTOALL){
     210           90 :         dtype = opParams.all2AllDataDes.sendType;
     211          560 :     } else if (opParams.opType == OpType::ALLTOALLV){
     212           73 :         dtype = opParams.all2AllVDataDes.sendType;
     213          487 :     } else if (opParams.opType == OpType::ALLTOALLVC){
     214            2 :         dtype = opParams.all2AllVCDataDes.sendType;
     215              :     }
     216          650 :     return dtype;
     217              : }
     218              : 
     219           81 : static void ReportErrMsg(const CollOpParams &opParams, DataType dtype)
     220              : {
     221           81 :     RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
     222              :                     std::vector<std::string>({"CheckOpDataTypeByMap", "[" + opParams.opType.Describe() + "][" + dtype.Describe() + "]",
     223              :                         "[opType][dataType]", "please check DataType that is not supported"}));
     224          243 :     HCCL_ERROR("[OpParamsChecker::CheckOpDataTypeByMap] opType [%s] with not support data type [%s], please check input opParam.",
     225              :                     opParams.opType.Describe().c_str(), dtype.Describe().c_str());
     226           81 : }
     227              : 
     228          653 : HcclResult OpParamsChecker::CheckOpDataTypeByMap(const CollOpParams &opParams, const DataTypeSupportMap &opData2TypeMap)
     229              : {
     230          653 :     auto iter = opData2TypeMap.find(opParams.opType);
     231          653 :     if (iter == opData2TypeMap.end()) {
     232            3 :         RPT_INPUT_ERR(true, "EI0003", std::vector<std::string>({"ccl_op", "value", "parameter", "expect"}),
     233              :             std::vector<std::string>({"CheckOpDataTypeByMap", opParams.opType.Describe(), "opType",
     234              :             "please check opType that is not supported"}));
     235            9 :         HCCL_ERROR("[OpParamsChecker::%s] invalid opType [%s], please check input opParam.",
     236              :                     __func__, opParams.opType.Describe().c_str());
     237            3 :         return HcclResult::HCCL_E_PARA;
     238              :     }
     239          650 :     bool checkResult = false;
     240          650 :     DataType dtype = GetDataType(opParams);
     241              : 
     242          650 :     if (opParams.opType == OpType::BATCHSENDRECV) {
     243           17 :         HcclSendRecvItem *sendRecvItems = static_cast<HcclSendRecvItem *>(opParams.batchSendRecvDataDes.sendRecvItemsPtr);
     244           17 :         u32 itemNum = opParams.batchSendRecvDataDes.itemNum;
     245              : 
     246           33 :         for (u32 i = 0; i < itemNum; ++i) {
     247           17 :             dtype = HcclDataTypeToDataType((sendRecvItems + i)->dataType);
     248           17 :             checkResult = (iter->second).test(static_cast<int>(dtype));
     249           17 :             if (!checkResult){
     250            1 :                 ReportErrMsg(opParams, dtype);
     251            1 :                 return HcclResult::HCCL_E_PARA;
     252              :             }
     253              :         }
     254              :     } else {
     255          633 :         checkResult = (iter->second).test(static_cast<int>(dtype));
     256          633 :         if (!checkResult){
     257           80 :             ReportErrMsg(opParams, dtype);
     258           80 :             return HcclResult::HCCL_E_PARA;
     259              :         }
     260              :     }
     261          569 :     return HcclResult::HCCL_SUCCESS;
     262            0 : }
     263              : 
     264              : DataTypeBitmap OpParamsChecker::dataTypeWithReduceAiv = DataTypeBitmap{}
     265              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     266              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     267              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     268              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT64))
     269              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     270              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     271              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16));
     272              : 
     273              : DataTypeBitmap OpParamsChecker::dataTypeWithoutReduceAiv = DataTypeBitmap{}
     274              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT8))
     275              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT16))
     276              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT32))
     277              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT64))
     278              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     279              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     280              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     281              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT64))
     282              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     283              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     284              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP64))
     285              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16))
     286              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_HIF8))
     287              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E4M3))
     288              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E5M2))
     289              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E8M0));
     290              : 
     291              : DataTypeBitmap OpParamsChecker::dataTypeWithReduceCcu = DataTypeBitmap{}
     292              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     293              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     294              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     295              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     296              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     297              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16));
     298              : DataTypeBitmap OpParamsChecker::dataTypeWithReduceAicpu = DataTypeBitmap{}
     299              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     300              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     301              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     302              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     303              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     304              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16))
     305              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP64))
     306              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT64))
     307              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT64));
     308              : DataTypeBitmap OpParamsChecker::dataTypeWithoutReduce = DataTypeBitmap{}
     309              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     310              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     311              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     312              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT64))
     313              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT8))
     314              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT16))
     315              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT32))
     316              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT64))
     317              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     318              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     319              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP64))
     320              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16))
     321              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_HIF8))
     322              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E4M3))
     323              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E5M2))
     324              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E8M0));
     325              : DataTypeBitmap OpParamsChecker::dataTypeWithoutReduceCcuOpbase = DataTypeBitmap{}
     326              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     327              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     328              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     329              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT64))
     330              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT8))
     331              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT16))
     332              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT32))
     333              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_UINT64))
     334              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     335              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     336              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP64))
     337              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16))
     338              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_HIF8))
     339              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E4M3))
     340              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E5M2))
     341              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E8M0));
     342              : DataTypeBitmap OpParamsChecker::dataTypeWithoutReduceCcuOffload = OpParamsChecker::dataTypeWithoutReduceCcuOpbase;
     343              : DataTypeBitmap OpParamsChecker::dataTypeWithReduceHost = DataTypeBitmap{}
     344              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     345              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT16))
     346              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT32))
     347              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     348              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     349              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16));
     350              : 
     351              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapAivOpbase = {
     352              :     {OpType::REDUCESCATTER, dataTypeWithReduceAiv},
     353              :     {OpType::ALLREDUCE, dataTypeWithReduceAiv},
     354              :     {OpType::ALLGATHER, dataTypeWithoutReduceAiv},
     355              :     {OpType::SCATTER, dataTypeWithoutReduceAiv},
     356              :     {OpType::ALLTOALL, dataTypeWithoutReduceAiv},
     357              :     {OpType::ALLTOALLV, dataTypeWithoutReduceAiv},
     358              :     {OpType::REDUCE, dataTypeWithReduceAiv},
     359              :     {OpType::BROADCAST, dataTypeWithoutReduceAiv},
     360              :     {OpType::SEND, dataTypeWithoutReduceAiv},
     361              :     {OpType::RECV, dataTypeWithoutReduceAiv},
     362              :     {OpType::BATCHSENDRECV, dataTypeWithoutReduceAiv}
     363              : };
     364              : 
     365              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapAivOffload= {
     366              :     {OpType::REDUCESCATTER, dataTypeWithReduceAiv},
     367              :     {OpType::ALLREDUCE, dataTypeWithReduceAiv},
     368              :     {OpType::ALLGATHER, dataTypeWithoutReduceAiv},
     369              :     {OpType::SCATTER, dataTypeWithoutReduceAiv},
     370              :     {OpType::ALLTOALL, dataTypeWithoutReduceAiv},
     371              :     {OpType::ALLTOALLV, dataTypeWithoutReduceAiv},
     372              :     {OpType::REDUCE, dataTypeWithReduceAiv},
     373              :     {OpType::BROADCAST, dataTypeWithoutReduceAiv}
     374              : };
     375              : 
     376              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapCcuOpbase = {
     377              :     {OpType::REDUCESCATTER, dataTypeWithReduceCcu},
     378              :     {OpType::ALLREDUCE, dataTypeWithReduceCcu},
     379              :     {OpType::ALLGATHER, dataTypeWithoutReduceCcuOpbase},
     380              :     {OpType::SCATTER, dataTypeWithoutReduce},
     381              :     {OpType::ALLTOALL, dataTypeWithoutReduceCcuOpbase},
     382              :     {OpType::ALLTOALLV, dataTypeWithoutReduceCcuOpbase},
     383              :     {OpType::REDUCE, dataTypeWithReduceCcu},
     384              :     {OpType::BROADCAST, dataTypeWithoutReduce},
     385              :     {OpType::REDUCESCATTERV, dataTypeWithReduceCcu},
     386              :     {OpType::ALLGATHERV, dataTypeWithoutReduceCcuOpbase}
     387              : };
     388              : 
     389              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapCcuOffload = {
     390              :     {OpType::REDUCESCATTER, dataTypeWithReduceCcu},
     391              :     {OpType::ALLREDUCE, dataTypeWithReduceCcu},
     392              :     {OpType::ALLGATHER, dataTypeWithoutReduceCcuOffload},
     393              :     {OpType::ALLTOALL, dataTypeWithoutReduce},
     394              :     {OpType::ALLTOALLV, dataTypeWithoutReduce},
     395              :     {OpType::REDUCE, dataTypeWithReduceCcu},
     396              :     {OpType::BROADCAST, dataTypeWithoutReduce},
     397              :     {OpType::REDUCESCATTERV, dataTypeWithReduceCcu},
     398              :     {OpType::ALLGATHERV, dataTypeWithoutReduceCcuOffload}
     399              : };
     400              : 
     401              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapAicpuOpbase = {
     402              :     {OpType::REDUCESCATTER, dataTypeWithReduceAicpu},
     403              :     {OpType::ALLREDUCE, dataTypeWithReduceAicpu},
     404              :     {OpType::ALLGATHER, dataTypeWithoutReduce},
     405              :     {OpType::SCATTER, dataTypeWithoutReduce},
     406              :     {OpType::ALLTOALL, dataTypeWithoutReduce},
     407              :     {OpType::ALLTOALLV, dataTypeWithoutReduce},
     408              :     {OpType::ALLTOALLVC, dataTypeWithoutReduce},
     409              :     {OpType::SEND, dataTypeWithoutReduce},
     410              :     {OpType::RECV, dataTypeWithoutReduce},
     411              :     {OpType::REDUCE, dataTypeWithReduceAicpu},
     412              :     {OpType::BROADCAST, dataTypeWithoutReduce},
     413              :     {OpType::BATCHSENDRECV, dataTypeWithoutReduce},
     414              :     {OpType::BATCHGET, dataTypeWithoutReduce},
     415              :     {OpType::BATCHPUT, dataTypeWithoutReduce}
     416              : };
     417              : 
     418              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapAicpuOffload = {
     419              :     {OpType::ALLGATHER, dataTypeWithoutReduce},
     420              :     {OpType::REDUCESCATTER, dataTypeWithReduceAicpu},
     421              :     {OpType::ALLREDUCE, dataTypeWithReduceAicpu},
     422              :     {OpType::ALLTOALL, dataTypeWithoutReduce},
     423              :     {OpType::ALLTOALLV, dataTypeWithoutReduce},
     424              :     {OpType::ALLTOALLVC, dataTypeWithoutReduce},
     425              :     {OpType::REDUCE, dataTypeWithReduceAicpu},
     426              :     {OpType::BROADCAST, dataTypeWithoutReduce},
     427              :     {OpType::SEND, dataTypeWithoutReduce},
     428              :     {OpType::RECV, dataTypeWithoutReduce}
     429              : };
     430              : 
     431              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapHostOffload = {
     432              :     {OpType::ALLGATHER, dataTypeWithoutReduce},
     433              :     {OpType::REDUCESCATTER, dataTypeWithReduceHost},
     434              :     {OpType::ALLREDUCE, dataTypeWithReduceHost},
     435              :     {OpType::ALLTOALL, dataTypeWithoutReduce},
     436              :     {OpType::ALLTOALLV, dataTypeWithoutReduce},
     437              :     {OpType::ALLTOALLVC, dataTypeWithoutReduce},
     438              :     {OpType::BROADCAST, dataTypeWithoutReduce},
     439              :     {OpType::SEND, dataTypeWithoutReduce},
     440              :     {OpType::RECV, dataTypeWithoutReduce},
     441              : };
     442              : 
     443              : DataTypeBitmap OpParamsChecker::dataTypeMC2HighP = dataTypeWithReduceCcu;
     444              : DataTypeBitmap OpParamsChecker::inputDataTypeMC2LowP = DataTypeBitmap{}
     445              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_INT8))
     446              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E5M2))
     447              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP8E4M3))
     448              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_HIF8));
     449              : DataTypeBitmap OpParamsChecker::OutputDataTypeMC2LowP = DataTypeBitmap{}
     450              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP16))
     451              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_FP32))
     452              :                                     | DataTypeBitmap(1 << static_cast<int>(HcclDataType::HCCL_DATA_TYPE_BFP16));
     453              : 
     454              : DataTypeSupportMap OpParamsChecker::opDataTypeSupportMapMC2 = {
     455              :     {OpType::ALLGATHER, dataTypeWithoutReduce},
     456              :     {OpType::REDUCESCATTER, dataTypeMC2HighP},
     457              :     {OpType::ALLREDUCE, dataTypeMC2HighP},
     458              :     {OpType::ALLTOALL, dataTypeWithoutReduce},
     459              :     {OpType::ALLTOALLV, dataTypeWithoutReduce},
     460              :     {OpType::HALFALLTOALLV, dataTypeWithoutReduce}
     461              : };
     462              : 
     463              : }
        

Generated by: LCOV version 2.0-1