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
|