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

Generated by: LCOV version 2.0-1