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 "coll_operator_check.h"
12 : #include "exception_util.h"
13 : #include "adapter_error_manager_pub.h"
14 :
15 : namespace Hccl {
16 :
17 13 : void ReportOpCheckFailed(
18 : const std::string& paraName, const std::string& localPara, const std::string& remotePara, const OpType& optype,
19 : const std::string& optag)
20 : {
21 13 : std::string opInfo = "Unknown";
22 38 : for (const auto& pair : HCOM_OP_TYPE_STR_MAP_V2) {
23 38 : if (pair.second == optype) {
24 13 : opInfo = std::string(pair.first);
25 13 : break;
26 : }
27 : }
28 : // 上报故障码EI0005
29 247 : RPT_INPUT_ERR(
30 : true, "EI0005", std::vector<std::string>({"ccl_op", "group", "para_name", "local_para", "remote_para"}),
31 : std::vector<std::string>({opInfo, optag, paraName, localPara, remotePara}));
32 26 : THROW<InvalidParamsException>(StringFormat(
33 : "[RankConsistentImpl][CompareFrame][%s]op information%s group%s %s check fail. "
34 : "local[%s], remote[%s]",
35 : __func__, opInfo.c_str(), optag.c_str(), paraName.c_str(), localPara.c_str(), remotePara.c_str()));
36 26 : }
37 :
38 9 : void ReportOpCheckFailed(
39 : const std::string& paraName, uint32_t localPara, uint32_t remotePara, const OpType& optype,
40 : const std::string& optag)
41 : {
42 9 : std::string opInfo = "Unknown";
43 34 : for (const auto& pair : HCOM_OP_TYPE_STR_MAP_V2) {
44 34 : if (pair.second == optype) {
45 9 : opInfo = std::string(pair.first);
46 9 : break;
47 : }
48 : }
49 : // 上报故障码EI0005
50 171 : RPT_INPUT_ERR(
51 : true, "EI0005", std::vector<std::string>({"ccl_op", "group", "para_name", "local_para", "remote_para"}),
52 : std::vector<std::string>({opInfo, optag, paraName, std::to_string(localPara), std::to_string(remotePara)}));
53 18 : THROW<InvalidParamsException>(StringFormat(
54 : "[RankConsistentImpl][CompareFrame][%s]op information%s group%s %s check fail. "
55 : "local[%u], remote[%u]",
56 : __func__, opInfo.c_str(), optag.c_str(), paraName.c_str(), localPara, remotePara));
57 18 : }
58 :
59 6 : void CompareDataDesOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
60 : {
61 6 : if (localOpData.dataDes.dataCount != remoteOpData.dataDes.dataCount) {
62 2 : ReportOpCheckFailed(
63 1 : "dataDes.dataCount", localOpData.dataDes.dataCount, remoteOpData.dataDes.dataCount, localOpData.opType,
64 1 : localOpData.opTag);
65 : }
66 :
67 5 : if (localOpData.dataDes.dataType != remoteOpData.dataDes.dataType) {
68 2 : ReportOpCheckFailed(
69 4 : "dataDes.dataType", localOpData.dataDes.dataType.Describe(), remoteOpData.dataDes.dataType.Describe(),
70 1 : localOpData.opType, localOpData.opTag);
71 : }
72 :
73 4 : if (localOpData.dataDes.strideCount != remoteOpData.dataDes.strideCount) {
74 2 : ReportOpCheckFailed(
75 1 : "dataDes.strideCount", localOpData.dataDes.strideCount, remoteOpData.dataDes.strideCount,
76 1 : localOpData.opType, localOpData.opTag);
77 : }
78 3 : }
79 :
80 2 : void CompareVDataDesOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
81 : {
82 2 : if (localOpData.vDataDes.dataType != remoteOpData.vDataDes.dataType) {
83 2 : ReportOpCheckFailed(
84 3 : "vDataDes.dataType", localOpData.vDataDes.dataType.Describe(), remoteOpData.vDataDes.dataType.Describe(),
85 1 : localOpData.opType, localOpData.opTag);
86 : }
87 1 : }
88 :
89 5 : void CompareAlltoAllOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
90 : {
91 5 : if (localOpData.all2AllDataDes.sendType != remoteOpData.all2AllDataDes.recvType) {
92 2 : ReportOpCheckFailed(
93 3 : "all2AllDataDes.sendType", localOpData.all2AllDataDes.sendType.Describe(),
94 3 : remoteOpData.all2AllDataDes.recvType.Describe(), localOpData.opType, localOpData.opTag);
95 : }
96 :
97 4 : if (localOpData.all2AllDataDes.recvType != remoteOpData.all2AllDataDes.sendType) {
98 2 : ReportOpCheckFailed(
99 3 : "all2AllDataDes.recvType", localOpData.all2AllDataDes.recvType.Describe(),
100 3 : remoteOpData.all2AllDataDes.sendType.Describe(), localOpData.opType, localOpData.opTag);
101 : }
102 :
103 3 : if (localOpData.all2AllDataDes.sendCount != remoteOpData.all2AllDataDes.recvCount) {
104 2 : ReportOpCheckFailed(
105 1 : "all2AllDataDes.sendCount", localOpData.all2AllDataDes.sendCount, remoteOpData.all2AllDataDes.recvCount,
106 1 : localOpData.opType, localOpData.opTag);
107 : }
108 :
109 2 : if (localOpData.all2AllDataDes.recvCount != remoteOpData.all2AllDataDes.sendCount) {
110 2 : ReportOpCheckFailed(
111 1 : "all2AllDataDes.recvCount", localOpData.all2AllDataDes.recvCount, remoteOpData.all2AllDataDes.sendCount,
112 1 : localOpData.opType, localOpData.opTag);
113 : }
114 1 : }
115 :
116 3 : void CompareAlltoAllVOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
117 : {
118 3 : if (localOpData.all2AllVDataDes.sendType != remoteOpData.all2AllVDataDes.recvType) {
119 2 : ReportOpCheckFailed(
120 3 : "all2AllVDataDes.sendType", localOpData.all2AllVDataDes.sendType.Describe(),
121 3 : remoteOpData.all2AllVDataDes.recvType.Describe(), localOpData.opType, localOpData.opTag);
122 : }
123 :
124 2 : if (localOpData.all2AllVDataDes.recvType != remoteOpData.all2AllVDataDes.sendType) {
125 2 : ReportOpCheckFailed(
126 3 : "all2AllVDataDes.recvType", localOpData.all2AllVDataDes.recvType.Describe(),
127 2 : remoteOpData.all2AllVDataDes.sendType.Describe(), localOpData.opType, localOpData.opTag);
128 : }
129 1 : }
130 :
131 3 : void CompareAlltoAllVCOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
132 : {
133 3 : if (localOpData.all2AllVCDataDes.sendType != remoteOpData.all2AllVCDataDes.recvType) {
134 2 : ReportOpCheckFailed(
135 3 : "all2AllVCDataDes.sendType", localOpData.all2AllVCDataDes.sendType.Describe(),
136 3 : remoteOpData.all2AllVCDataDes.recvType.Describe(), localOpData.opType, localOpData.opTag);
137 : }
138 :
139 2 : if (localOpData.all2AllVCDataDes.recvType != remoteOpData.all2AllVCDataDes.sendType) {
140 2 : ReportOpCheckFailed(
141 3 : "all2AllVCDataDes.recvType", localOpData.all2AllVCDataDes.recvType.Describe(),
142 2 : remoteOpData.all2AllVCDataDes.sendType.Describe(), localOpData.opType, localOpData.opTag);
143 : }
144 1 : }
145 :
146 29 : void CompareNormalOp(const CollOperator& localOpData, const CollOperator& remoteOpData)
147 : {
148 29 : if (localOpData.opMode != remoteOpData.opMode) {
149 0 : ReportOpCheckFailed(
150 0 : "opMode", localOpData.opMode.Describe(), remoteOpData.opMode.Describe(), localOpData.opType,
151 0 : localOpData.opTag);
152 : }
153 :
154 29 : if (localOpData.opType == OpType::SEND) {
155 2 : if (remoteOpData.opType != OpType::RECV) {
156 0 : ReportOpCheckFailed(
157 0 : "opType", localOpData.opType.Describe(), remoteOpData.opType.Describe(), localOpData.opType,
158 0 : localOpData.opTag);
159 : }
160 27 : } else if (localOpData.opType == OpType::RECV) {
161 0 : if (remoteOpData.opType != OpType::SEND) {
162 0 : ReportOpCheckFailed(
163 0 : "opType", localOpData.opType.Describe(), remoteOpData.opType.Describe(), localOpData.opType,
164 0 : localOpData.opTag);
165 : }
166 27 : } else if (localOpData.opType != remoteOpData.opType) {
167 2 : ReportOpCheckFailed(
168 4 : "opType", localOpData.opType.Describe(), remoteOpData.opType.Describe(), localOpData.opType,
169 1 : localOpData.opTag);
170 : }
171 :
172 28 : if (localOpData.reduceOp != remoteOpData.reduceOp) {
173 2 : ReportOpCheckFailed(
174 4 : "reduceOp", localOpData.reduceOp.Describe(), remoteOpData.reduceOp.Describe(), localOpData.opType,
175 1 : localOpData.opTag);
176 : }
177 :
178 27 : if (localOpData.dataType != remoteOpData.dataType) {
179 2 : ReportOpCheckFailed(
180 4 : "dataType", localOpData.dataType.Describe(), remoteOpData.dataType.Describe(), localOpData.opType,
181 1 : localOpData.opTag);
182 : }
183 :
184 26 : if (localOpData.opType != OpType::ALLGATHERV && localOpData.opType != OpType::REDUCESCATTERV) {
185 24 : if (localOpData.dataCount != remoteOpData.dataCount) {
186 2 : ReportOpCheckFailed(
187 1 : "dataCount", localOpData.dataCount, remoteOpData.dataCount, localOpData.opType, localOpData.opTag);
188 : }
189 : }
190 :
191 25 : if (localOpData.root != remoteOpData.root) {
192 3 : ReportOpCheckFailed("root", localOpData.root, remoteOpData.root, localOpData.opType, localOpData.opTag);
193 : }
194 :
195 24 : if (localOpData.opType == OpType::SEND || localOpData.opType == OpType::RECV) {
196 2 : if (localOpData.myRank != remoteOpData.sendRecvRemoteRank) {
197 2 : ReportOpCheckFailed(
198 1 : "sendRecvRemoteRank", localOpData.myRank, remoteOpData.sendRecvRemoteRank, localOpData.opType,
199 1 : localOpData.opTag);
200 : }
201 : }
202 :
203 23 : if (localOpData.opTag != remoteOpData.opTag) {
204 3 : ReportOpCheckFailed("opTag", localOpData.opTag, remoteOpData.opTag, localOpData.opType, localOpData.opTag);
205 : }
206 :
207 22 : if (localOpData.staticAddr != remoteOpData.staticAddr) {
208 2 : ReportOpCheckFailed(
209 1 : "staticAddr", static_cast<uint32_t>(localOpData.staticAddr), static_cast<uint32_t>(remoteOpData.staticAddr),
210 1 : localOpData.opType, localOpData.opTag);
211 : }
212 :
213 21 : if (localOpData.staticShape != remoteOpData.staticShape) {
214 2 : ReportOpCheckFailed(
215 1 : "staticShape", static_cast<uint32_t>(localOpData.staticShape),
216 1 : static_cast<uint32_t>(remoteOpData.staticShape), localOpData.opType, localOpData.opTag);
217 : }
218 :
219 20 : if (localOpData.outputDataType != remoteOpData.outputDataType) {
220 2 : ReportOpCheckFailed(
221 3 : "outputDataType", localOpData.outputDataType.Describe(), remoteOpData.outputDataType.Describe(),
222 1 : localOpData.opType, localOpData.opTag);
223 : }
224 19 : }
225 :
226 : /*
227 : 当前校验类型不支持vDataDes和all2AllVCDataDes相关内容;batchSendRecvDataDes中的itemNum字段没有校验的必要,双端该值可能不相等
228 : */
229 29 : void CheckCollOperator(const CollOperator& localOpData, const CollOperator& remoteOpData)
230 : {
231 29 : CompareNormalOp(localOpData, remoteOpData);
232 :
233 19 : if (localOpData.opType == OpType::BATCHSENDRECV) {
234 0 : return;
235 : }
236 :
237 19 : if (localOpData.opType == OpType::ALLTOALL) {
238 5 : CompareAlltoAllOp(localOpData, remoteOpData);
239 1 : return;
240 : }
241 :
242 14 : if (localOpData.opType == OpType::ALLTOALLV) {
243 3 : CompareAlltoAllVOp(localOpData, remoteOpData);
244 1 : return;
245 : }
246 :
247 11 : if (localOpData.opType == OpType::ALLTOALLVC) {
248 3 : CompareAlltoAllVCOp(localOpData, remoteOpData);
249 1 : return;
250 : }
251 :
252 8 : if (localOpData.opType == OpType::ALLGATHERV || localOpData.opType == OpType::REDUCESCATTERV) {
253 2 : CompareVDataDesOp(localOpData, remoteOpData);
254 1 : return;
255 : }
256 :
257 6 : CompareDataDesOp(localOpData, remoteOpData);
258 3 : return;
259 : }
260 : } // namespace Hccl
|