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 : }
|