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 "executor_utils.h"
12 :
13 : namespace Hccl {
14 1 : bool IsEnableCounterNotifyByDevType(const RankId myRank, const DevType devType)
15 : {
16 1 : switch (devType) {
17 1 : case DevType::DEV_TYPE_950:
18 : case DevType::DEV_TYPE_960:
19 3 : HCCL_DEBUG("[CollAlgFactory] Rank [%d], CounterNotify func enabled.", myRank);
20 1 : return true;
21 0 : default:
22 0 : HCCL_DEBUG("[CollAlgFactory] Rank [%d], CounterNotify func disabled.", myRank);
23 0 : return false;
24 : }
25 : }
26 :
27 1 : HcclResult InitOpInfo(const CollAlgOperator &op, OpType &opType, ReduceOp &redOp, u32 &root)
28 : {
29 1 : opType = op.opType;
30 1 : switch (opType) {
31 0 : case OpType::ALLREDUCE:
32 : case OpType::REDUCESCATTER:
33 0 : redOp = op.reduceOp;
34 0 : break;
35 0 : case OpType::SCATTER:
36 : case OpType::BROADCAST:
37 0 : root = op.root;
38 0 : break;
39 0 : case OpType::REDUCE:
40 0 : redOp = op.reduceOp;
41 0 : root = op.root;
42 0 : break;
43 1 : default:
44 1 : break;
45 : }
46 1 : return HcclResult::HCCL_SUCCESS;
47 : }
48 :
49 0 : HcclResult InitDataInfo(const CollAlgOperator &op, DataType &dataType, DataType &outputDataType, u64 &dataCount)
50 : {
51 0 : dataType = op.dataType;
52 0 : dataCount = op.dataCount;
53 0 : outputDataType = op.outputDataType;
54 0 : if (outputDataType == DataType::INVALID) {
55 0 : outputDataType = dataType;
56 : }
57 :
58 0 : return HcclResult::HCCL_SUCCESS;
59 : }
60 :
61 : // Get Prior Link from virtual topo
62 4 : const std::vector<NetInstance::Path> GetPathsFromRankGraph(
63 : const RankGraph *rankGraph, const RankId srcRank, const RankId dstRank)
64 : {
65 : // 遍历当前节点的所有层级,返回两个节点间查到到的所有path
66 4 : std::vector<NetInstance::Path> pathList;
67 4 : std::set<u32> levelSet = rankGraph->GetLevels(srcRank);
68 8 : for (u32 levelIdx : levelSet) {
69 4 : std::vector<NetInstance::Path> paths = rankGraph->GetPaths(levelIdx, srcRank, dstRank);
70 4 : pathList.insert(pathList.end(), paths.begin(), paths.end());
71 4 : }
72 4 : return pathList;
73 4 : }
74 :
75 2 : HcclResult AddToResLinks(const RankId vNeighborRank, const LinkData &linkData, ResLinks &resLinks)
76 : {
77 6 : HCCL_DEBUG(
78 : "RankId [%d] linkData.des[%s] resLinks[%zu]", vNeighborRank, linkData.Describe().c_str(), resLinks.size());
79 2 : auto rankLinkIter = resLinks.find(vNeighborRank);
80 2 : if (rankLinkIter == resLinks.end()) {
81 4 : std::vector<LinkData> tmpLinks = {linkData};
82 2 : resLinks.insert(std::pair<RankId, std::vector<LinkData>>(vNeighborRank, tmpLinks));
83 2 : } else {
84 0 : rankLinkIter->second.push_back(linkData);
85 : }
86 2 : return HcclResult::HCCL_SUCCESS;
87 : }
88 :
89 1 : HcclResult PrepResLinks(const RankId myRank, const RankGraph *rankGraph, const std::vector<BasePortType> &linkPriority,
90 : const LinkReq &linkReq, ResLinks &resLinks)
91 : {
92 3 : HCCL_DEBUG("PrepResLinks linkPriority.size()[%zu], linkReq.size()[%zu]", linkPriority.size(), linkReq.size());
93 3 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
94 2 : const std::vector<NetInstance::Path> tmpPaths = GetPathsFromRankGraph(rankGraph, myRank, resReqIter->first);
95 2 : if (resReqIter->second == 1) {
96 2 : CHK_PRT_RET(tmpPaths.size() == 0,
97 : HCCL_ERROR("[CollAlgFactory] Unable to obtain valid link, srcRank [%d], dstRank [%d].", myRank,
98 : resReqIter->first),
99 : HcclResult::HCCL_E_INTERNAL);
100 2 : LinkData requiredLinkData(tmpPaths[0]); // 当前只取第一条path
101 : // updata res
102 2 : CHK_PRT_RET(AddToResLinks(resReqIter->first, requiredLinkData, resLinks) != HcclResult::HCCL_SUCCESS,
103 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Fail to prepare links.", myRank), HcclResult::HCCL_E_INTERNAL);
104 : } else {
105 0 : CHK_PRT_RET(tmpPaths.size() < resReqIter->second,
106 : HCCL_ERROR("[CollAlgFactory] Rank [%d], available linkNum smaller than required.", myRank),
107 : HcclResult::HCCL_E_INTERNAL);
108 : // 从所有path中选择前resReqIter->second条
109 0 : for (u32 linkNum = 0; linkNum < resReqIter->second; linkNum++) {
110 0 : LinkData requiredLinkData(tmpPaths[linkNum]);
111 : // updata res
112 0 : CHK_PRT_RET(AddToResLinks(resReqIter->first, requiredLinkData, resLinks) != HcclResult::HCCL_SUCCESS,
113 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Fail to prepare links.", myRank),
114 : HcclResult::HCCL_E_INTERNAL);
115 : }
116 : }
117 2 : }
118 :
119 1 : return HcclResult::HCCL_SUCCESS;
120 : }
121 :
122 0 : HcclResult PrepResLinks(const RankId myRank, const LinkReq &linkReq, ConnectedLinkMgr *linkMgr, ResLinks &resLinks)
123 : {
124 0 : CHK_PTR_NULL(linkMgr);
125 0 : HCCL_DEBUG("PrepResLinks linkReq.size()[%zu]", linkReq.size());
126 0 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
127 0 : if (resReqIter->second == 1) {
128 0 : auto rankId = resReqIter->first;
129 0 : auto links = linkMgr->GetLinks(rankId);
130 0 : CHK_PRT_RET(links.size() == 0, HCCL_ERROR("[PrepResLinks] Rank [%d], Fail to get peer links.", myRank),
131 : HcclResult::HCCL_E_INTERNAL);
132 0 : LinkData requiredLinkData = links[0];
133 : // updata res
134 0 : CHK_PRT_RET(AddToResLinks(resReqIter->first, requiredLinkData, resLinks) != HcclResult::HCCL_SUCCESS,
135 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Fail to prepare links.", myRank),
136 : HcclResult::HCCL_E_INTERNAL);
137 0 : } else {
138 0 : for (u32 linkNum = 0; linkNum < resReqIter->second; linkNum++) {
139 0 : LinkData requiredLinkData = linkMgr->GetLinks(resReqIter->first)[linkNum];
140 : // updata res
141 0 : CHK_PRT_RET(AddToResLinks(resReqIter->first, requiredLinkData, resLinks) != HcclResult::HCCL_SUCCESS,
142 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Fail to prepare links.", myRank),
143 : HcclResult::HCCL_E_INTERNAL);
144 : }
145 : }
146 : }
147 0 : return HcclResult::HCCL_SUCCESS;
148 : }
149 :
150 1 : HcclResult CalcResLinks(const RankId myRank, const RankGraph *rankGraph, const std::vector<BasePortType> &linkPriority,
151 : const LinkReq &linkReq, std::vector<LinkData> &links)
152 : {
153 3 : HCCL_DEBUG("CalcResLinks linkPriority.size()[%zu]", linkPriority.size());
154 3 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
155 2 : const std::vector<NetInstance::Path> tmpPaths = GetPathsFromRankGraph(rankGraph, myRank, resReqIter->first);
156 2 : if (resReqIter->second == 1) {
157 2 : CHK_PRT_RET(tmpPaths.size() == 0,
158 : HCCL_ERROR("[CollAlgFactory] Unable to obtain valid link, srcRank [%d], dstRank [%d].", myRank,
159 : resReqIter->first),
160 : HcclResult::HCCL_E_INTERNAL);
161 : // updata res
162 2 : links.emplace_back(tmpPaths[0]);
163 : } else {
164 0 : CHK_PRT_RET(tmpPaths.size() < resReqIter->second,
165 : HCCL_ERROR("[CollAlgFactory] Rank [%d], available linkNum smaller than required.", myRank),
166 : HcclResult::HCCL_E_INTERNAL);
167 0 : for (u32 linkNum = 0; linkNum < resReqIter->second; linkNum++) {
168 : // updata res
169 0 : links.emplace_back(tmpPaths[linkNum]);
170 : }
171 : }
172 2 : }
173 :
174 1 : return HcclResult::HCCL_SUCCESS;
175 : }
176 :
177 1 : HcclResult CalcLinkInfo(const RankId myRank, const RankGraph *rankGraph, const LinkReq &linkReq,
178 : std::vector<std::pair<u32, RankId>> &algTempLinksInfo)
179 : {
180 1 : std::set<u32> levelSet = rankGraph->GetLevels(myRank);
181 3 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
182 2 : RankId remoteRank = resReqIter->first;
183 2 : if (resReqIter->second == 0) {
184 2 : continue;
185 : }
186 2 : if (levelSet.size() == 1) {
187 2 : algTempLinksInfo.push_back(std::make_pair(0, remoteRank));
188 2 : continue;
189 : }
190 : // 当前场景只考虑两层拓扑场景
191 0 : u32 levelIdx = 0;
192 0 : const NetInstance *netInstance = rankGraph->GetNetInstanceByRankId(levelIdx, myRank);
193 0 : std::set<RankId> rankSet = netInstance->GetRankIds();
194 0 : auto rankInRankSet = std::find(rankSet.begin(), rankSet.end(), remoteRank);
195 0 : if (rankInRankSet != rankSet.end()) {
196 0 : algTempLinksInfo.push_back(std::make_pair(0, remoteRank));
197 : } else {
198 0 : algTempLinksInfo.push_back(std::make_pair(1, remoteRank));
199 : }
200 0 : }
201 1 : return HcclResult::HCCL_SUCCESS;
202 1 : }
203 :
204 0 : HcclResult SetPathNumMapByRankGraphMultiLevel(const RankGraph *rankGraph, std::vector<std::vector<RankId>>&virtRanks_,
205 : RankId myRank_, std::vector<map<u32, u32>>&rank2PathNumMap){
206 0 : uint64_t levelNum = 2;
207 0 : for (uint64_t levelNumIdx = 0; levelNumIdx < levelNum; levelNumIdx++) {
208 0 : rank2PathNumMap.emplace_back();
209 0 : for (auto rankIdx : virtRanks_[levelNumIdx]) {
210 0 : if (rankIdx == myRank_) {
211 0 : continue;
212 : }
213 0 : std::vector<NetInstance::Path> tmpPaths = rankGraph->GetPaths(levelNumIdx, myRank_, rankIdx);
214 0 : auto pathNum = 0;
215 0 : for (const auto &path : tmpPaths) {
216 0 : bool isWithPcie = false;
217 0 : for (const auto &link : path.links) {
218 0 : if (*link.GetLinkProtocols().begin() == LinkProtocol::PCIE) {
219 0 : isWithPcie = true;
220 0 : break;
221 : }
222 : }
223 0 : if (!isWithPcie) {
224 0 : pathNum++;
225 : }
226 : }
227 0 : rank2PathNumMap[levelNumIdx][rankIdx] = pathNum;
228 0 : HCCL_INFO("[%s]levelNumIdx[%u] rankIdx[%d] pathNum[%d]", __func__, levelNumIdx, rankIdx, pathNum);
229 0 : }
230 : }
231 0 : if(rank2PathNumMap.size() == 0){
232 0 : HCCL_ERROR("No path to all remoteRank");
233 0 : return HcclResult::HCCL_E_INTERNAL;
234 : }
235 0 : return HcclResult::HCCL_SUCCESS;
236 : }
237 :
238 0 : HcclResult SetPathNumMapByRankGraphMultiLevel(const RankGraph *rankGraph, std::vector<RankId>&virtRanks_,
239 : RankId myRank_, std::map<u32, u32>&rank2PathNumMap){
240 0 : std::set<u32> levelSet = rankGraph->GetLevels(myRank_);
241 0 : for(auto level : levelSet){
242 0 : bool levelFlag=1;
243 0 : for(auto rankIdx : virtRanks_){
244 0 : if(rankIdx == myRank_){
245 0 : continue;
246 : }
247 : std::vector<NetInstance::Path> tmpPaths =
248 0 : rankGraph->GetPaths(level, myRank_, rankIdx);
249 0 : if(tmpPaths.size()==0){
250 0 : rank2PathNumMap.clear();
251 0 : levelFlag = 0;
252 0 : break;
253 : }
254 0 : auto pathNum = 0;
255 0 : for (const auto &path : tmpPaths) {
256 0 : bool isWithPcie = false;
257 0 : for (const auto &link : path.links) {
258 0 : if (*link.GetLinkProtocols().begin() == LinkProtocol::PCIE) {
259 0 : isWithPcie = true;
260 0 : break;
261 : }
262 : }
263 0 : if (!isWithPcie) {
264 0 : pathNum++;
265 : }
266 : }
267 0 : rank2PathNumMap[rankIdx] = pathNum;
268 0 : HCCL_INFO("[%s]rankIdx[%d] pathNum[%d]", __func__, rankIdx, pathNum);
269 0 : }
270 0 : if(levelFlag){
271 0 : break;
272 : }
273 : }
274 0 : if(rank2PathNumMap.size() == 0){
275 0 : HCCL_ERROR("No path to all remoteRank");
276 0 : return HcclResult::HCCL_E_INTERNAL;
277 : }
278 0 : return HcclResult::HCCL_SUCCESS;
279 0 : }
280 :
281 0 : HcclResult SetPathNumMapByLinkMgrMultiLevel(ConnectedLinkMgr*linkMgr, std::vector<std::vector<RankId>>&virtRanks_,
282 : RankId myRank_, std::vector<map<u32, u32>>&rank2PathNumMap){
283 : (void) myRank_;
284 0 : uint64_t levelNum = 2;
285 0 : for (uint64_t levelNumIdx = 0; levelNumIdx < levelNum; levelNumIdx++) {
286 0 : rank2PathNumMap.emplace_back();
287 0 : for (auto rankIdx : virtRanks_[levelNumIdx]) {
288 0 : auto links = linkMgr->GetLinks(levelNumIdx, rankIdx);
289 0 : auto linkNum = 0;
290 0 : for (const auto& link : links) {
291 0 : if (link.GetLinkProtocol() != LinkProtocol::PCIE) {
292 0 : linkNum++;
293 : }
294 : }
295 0 : if (linkNum != 0){
296 0 : rank2PathNumMap[levelNumIdx][rankIdx] = linkNum;
297 : }
298 0 : HCCL_INFO("[%s]levelNumIdx[%u] rankIdx[%d] linkNum[%d]", __func__, levelNumIdx, rankIdx, linkNum);
299 0 : }
300 : }
301 0 : if(rank2PathNumMap.size() == 0){
302 0 : HCCL_ERROR("No path to all remoteRank");
303 0 : return HcclResult::HCCL_E_INTERNAL;
304 : }
305 0 : return HcclResult::HCCL_SUCCESS;
306 : }
307 :
308 0 : HcclResult SetPathNumMapByLinkMgrMultiLevel(ConnectedLinkMgr*linkMgr, std::vector<RankId>&virtRanks_,
309 : RankId myRank_, map<u32, u32>&rank2PathNumMap){
310 : (void) myRank_;
311 0 : for(u32 rankIdx:virtRanks_){
312 0 : auto links = linkMgr->GetLinks(rankIdx);
313 0 : auto linkNum = 0;
314 0 : for (const auto& link : links) {
315 0 : if (link.GetLinkProtocol() != LinkProtocol::PCIE) {
316 0 : linkNum++;
317 : }
318 : }
319 0 : if (linkNum != 0){
320 0 : rank2PathNumMap[rankIdx] = linkNum;
321 : }
322 0 : HCCL_INFO("[%s]rankIdx[%u] linkNum[%d]", __func__, rankIdx, linkNum);
323 0 : }
324 0 : if(rank2PathNumMap.size() == 0){
325 0 : HCCL_ERROR("No path to all remoteRank");
326 0 : return HcclResult::HCCL_E_INTERNAL;
327 : }
328 0 : return HcclResult::HCCL_SUCCESS;
329 : }
330 :
331 : } // namespace Hccl
|