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 "ins_alg_template_base.h"
12 : #include "log.h"
13 :
14 : namespace Hccl {
15 :
16 :
17 0 : InsAlgTemplateBase::InsAlgTemplateBase(const RankId virtualRank, const u32 tempRankSize,
18 : const std::vector<std::vector<RankId>> &tempVTopo,
19 0 : const std::map<RankId, u32> &tempVirtRankMap)
20 0 : : myRank_(virtualRank), tempRankSize_(tempRankSize), tempVTopo_(tempVTopo), tempVirtRankMap_(tempVirtRankMap)
21 : {
22 0 : }
23 :
24 0 : InsAlgTemplateBase::~InsAlgTemplateBase()
25 : {
26 0 : }
27 :
28 0 : void InsAlgTemplateBase::SetCollOp(const CollAlgOperator &op)
29 : {
30 0 : op_ = op;
31 0 : return;
32 : }
33 :
34 0 : void InsAlgTemplateBase::SetDmaMode(const DmaMode dmaMode)
35 : {
36 0 : dmaMode_ = dmaMode;
37 0 : return;
38 : }
39 :
40 0 : void InsAlgTemplateBase::SetRoot(const u32 root)
41 : {
42 0 : root_ = root;
43 0 : return;
44 : }
45 :
46 0 : u64 InsAlgTemplateBase::CalcLoopMaxCount(ParamPool ¶mPool)
47 : {
48 0 : u64 loopMaxCount = 0;
49 0 : if (paramPool.params.opMode == OpMode::OPBASE) {
50 0 : u64 maxLoopSize = std::min(static_cast<u64>(paramPool.params.maxTmpMemSize), static_cast<u64>(UB_MAX_DATA_SIZE));
51 0 : loopMaxCount = maxLoopSize / (DataTypeSizeGet(paramPool.op.dataType) * tempRankSize_) * tempRankSize_;
52 : } else {
53 0 : loopMaxCount = paramPool.op.dataCount;
54 : }
55 0 : return loopMaxCount;
56 : }
57 :
58 0 : HcclResult InsAlgTemplateBase::PostCopyOpbase(const UsrData &usrData, std::vector<InsQuePtr> &tempInsQues) const
59 : {
60 0 : for (size_t i = 0; i < usrData.scratchOutSlices.size(); i++) {
61 : std::unique_ptr<Instruction> insLocalCopy
62 0 : = std::make_unique<InsLocalCopy>(usrData.scratchOutSlices[i], usrData.usrOutSlices[i]);
63 0 : tempInsQues[0]->Append(std::move(insLocalCopy));
64 0 : }
65 :
66 0 : return HcclResult::HCCL_SUCCESS;
67 : }
68 :
69 0 : HcclResult InsAlgTemplateBase::PreCopyOpbase(const UsrData &usrData, std::vector<InsQuePtr> &tempInsQues) const
70 : {
71 0 : for (size_t i = 0; i < usrData.usrInSlices.size(); i++) {
72 : std::unique_ptr<Instruction> insLocalCopy
73 0 : = std::make_unique<InsLocalCopy>(usrData.usrInSlices[i], usrData.scratchInSlices[i]);
74 0 : tempInsQues[0]->Append(std::move(insLocalCopy));
75 0 : }
76 :
77 0 : return HcclResult::HCCL_SUCCESS;
78 : }
79 :
80 0 : HcclResult InsAlgTemplateBase::CalcSliceInfo(const AllignInfo &allignInfo, const u64 dataSize,
81 : RankSliceInfo &sliceInfoVec)
82 : {
83 : (void)allignInfo;
84 : (void)dataSize;
85 : (void)sliceInfoVec;
86 0 : HCCL_ERROR("[InsCollAlgFactory] Unsupported interface of slice info calculation!");
87 0 : return HcclResult::HCCL_E_INTERNAL;
88 : }
89 :
90 0 : HcclResult InsAlgTemplateBase::CalcRes(AlgTempResReq &tempResReq)
91 : {
92 : (void)tempResReq;
93 0 : HCCL_ERROR("[InsCollAlgFactory] Unsupported interface of resource calculation!");
94 0 : return HcclResult::HCCL_E_INTERNAL;
95 : }
96 :
97 0 : HcclResult InsAlgTemplateBase::CalcResDetour(const RankGraph *rankGraph, AlgTempResReq &tempResReq)
98 : {
99 : (void)rankGraph;
100 : (void)tempResReq;
101 0 : HCCL_ERROR("[InsCollAlgFactory] Current alg do not support detour mode!");
102 0 : return HcclResult::HCCL_E_INTERNAL;
103 : }
104 :
105 0 : HcclResult InsAlgTemplateBase::CalcResDetour(ConnectedLinkMgr *linkMgr, AlgTempResReq &tempResReq)
106 : {
107 : (void)linkMgr;
108 : (void)tempResReq;
109 0 : HCCL_ERROR("[InsCollAlgFactory] Current alg do not support detour mode!");
110 0 : return HcclResult::HCCL_E_INTERNAL;
111 : }
112 :
113 0 : uint64_t InsAlgTemplateBase::GetMaxSliceSize()
114 : {
115 0 : return UB_MAX_DATA_SIZE; // return max value
116 : }
117 :
118 0 : void InsAlgTemplateBase::InitReduceInfo(const ReduceOp &redOp, const DataType &dataType)
119 : {
120 0 : redOp_ = redOp;
121 0 : dataType_ = dataType;
122 0 : return;
123 : }
124 :
125 0 : HcclResult InsAlgTemplateBase::Run(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
126 : const BuffInfo &buffInfo, const ResLinks &tempLinks,
127 : std::vector<InsQuePtr> &tempInsQues)
128 : {
129 : (void)tempFuncs;
130 : (void)sliceInfoVec;
131 : (void)buffInfo;
132 : (void)tempLinks;
133 : (void)tempInsQues;
134 0 : HCCL_ERROR("[InsAlgTemplateBase] Unsupported interface of GenInsQue!");
135 0 : return HcclResult::HCCL_E_INTERNAL;
136 : }
137 :
138 0 : void InsAlgTemplateBase::SetDataType(const DataType &dataType)
139 : {
140 0 : dataType_ = dataType;
141 0 : return;
142 : }
143 :
144 0 : void InsAlgTemplateBase::SetReduceOp(const ReduceOp &redOp)
145 : {
146 0 : redOp_ = redOp;
147 0 : return;
148 : }
149 :
150 0 : HcclResult InsAlgTemplateBase::PreSync(const u32 queIdx, std::vector<InsQuePtr> &syncInsQues) const
151 : {
152 0 : InsQuePtr currInsQue = syncInsQues[queIdx];
153 0 : if (queIdx == 0) {
154 : // Semaphore Post
155 0 : if (enableCounterNotify_) {
156 0 : std::unique_ptr<InsLocalBcastPost> insLocalBcastPost = std::make_unique<InsLocalBcastPost>(0);
157 0 : for (size_t qidx = 1; qidx < syncInsQues.size(); qidx++) {
158 0 : insLocalBcastPost->Append(syncInsQues[qidx]->GetId());
159 : }
160 0 : CHK_PTR_NULL(insLocalBcastPost);
161 0 : currInsQue->Append(std::move(insLocalBcastPost));
162 0 : } else {
163 0 : for (size_t qidx = 1; qidx < syncInsQues.size(); qidx++) {
164 : std::unique_ptr<Instruction> insLocalPostTo
165 0 : = std::make_unique<InsLocalPostTo>(syncInsQues[qidx]->GetId());
166 0 : CHK_PTR_NULL(insLocalPostTo);
167 0 : currInsQue->Append(std::move(insLocalPostTo));
168 0 : }
169 : }
170 : } else {
171 : // Semaphore Wait
172 0 : if (enableCounterNotify_) {
173 : std::unique_ptr<Instruction> insLocalWaitFrom
174 0 : = std::make_unique<InsLocalWaitFrom>(syncInsQues[0]->GetId(), NotifyType::COUNTER);
175 0 : CHK_PTR_NULL(insLocalWaitFrom);
176 0 : currInsQue->Append(std::move(insLocalWaitFrom));
177 0 : } else {
178 0 : std::unique_ptr<Instruction> insLocalWaitFrom = std::make_unique<InsLocalWaitFrom>(syncInsQues[0]->GetId());
179 0 : CHK_PTR_NULL(insLocalWaitFrom);
180 0 : currInsQue->Append(std::move(insLocalWaitFrom));
181 0 : }
182 : }
183 :
184 0 : return HcclResult::HCCL_SUCCESS;
185 0 : }
186 :
187 0 : HcclResult InsAlgTemplateBase::PostSync(const u32 queIdx, std::vector<InsQuePtr> &syncInsQues) const
188 : {
189 0 : InsQuePtr currInsQue = syncInsQues[queIdx];
190 0 : if (queIdx == 0) {
191 : // Semaphore Wait
192 0 : if (enableCounterNotify_) {
193 0 : std::unique_ptr<InsLocalWaitGroup> insLocalWaitGroup = std::make_unique<InsLocalWaitGroup>(0);
194 0 : for (size_t qidx = 1; qidx < syncInsQues.size(); qidx++) {
195 0 : insLocalWaitGroup->Append(syncInsQues[qidx]->GetId());
196 : }
197 0 : CHK_PTR_NULL(insLocalWaitGroup);
198 0 : currInsQue->Append(std::move(insLocalWaitGroup));
199 0 : } else {
200 0 : for (size_t qidx = 1; qidx < syncInsQues.size(); qidx++) {
201 : std::unique_ptr<Instruction> insLocalWaitFrom
202 0 : = std::make_unique<InsLocalWaitFrom>(syncInsQues[qidx]->GetId());
203 0 : CHK_PTR_NULL(insLocalWaitFrom);
204 0 : currInsQue->Append(std::move(insLocalWaitFrom));
205 0 : }
206 : }
207 : } else {
208 : // Semaphore Post
209 0 : if (enableCounterNotify_) {
210 : std::unique_ptr<Instruction> insLocalPostTo
211 0 : = std::make_unique<InsLocalPostTo>(syncInsQues[0]->GetId(), NotifyType::COUNTER);
212 0 : CHK_PTR_NULL(insLocalPostTo);
213 0 : currInsQue->Append(std::move(insLocalPostTo));
214 0 : } else {
215 0 : std::unique_ptr<Instruction> insLocalPostTo = std::make_unique<InsLocalPostTo>(syncInsQues[0]->GetId());
216 0 : CHK_PTR_NULL(insLocalPostTo);
217 0 : currInsQue->Append(std::move(insLocalPostTo));
218 0 : }
219 : }
220 :
221 0 : return HcclResult::HCCL_SUCCESS;
222 0 : }
223 :
224 0 : HcclResult InsAlgTemplateBase::PreSyncInterQueues(std::vector<InsQuePtr> &syncInsQues) const
225 : {
226 0 : for (size_t queIdx = 0; queIdx < syncInsQues.size(); queIdx++) {
227 0 : CHK_PRT_RET(PreSync(queIdx, syncInsQues) != HcclResult::HCCL_SUCCESS,
228 : HCCL_ERROR("[InsCollAlgFactory] Rank [%d], Que [%u], Semaphore Synchronization Failed.", myRank_,
229 : syncInsQues[queIdx]->GetId()),
230 : HcclResult::HCCL_E_INTERNAL);
231 : }
232 :
233 0 : return HcclResult::HCCL_SUCCESS;
234 : }
235 :
236 0 : HcclResult InsAlgTemplateBase::PostSyncInterQueues(std::vector<InsQuePtr> &syncInsQues) const
237 : {
238 0 : for (size_t queIdx = 0; queIdx < syncInsQues.size(); queIdx++) {
239 0 : CHK_PRT_RET(PostSync(queIdx, syncInsQues) != HcclResult::HCCL_SUCCESS,
240 : HCCL_ERROR("[InsCollAlgFactory] Rank [%d], Que [%u], Semaphore Synchronization Failed.", myRank_,
241 : syncInsQues[queIdx]->GetId()),
242 : HcclResult::HCCL_E_INTERNAL);
243 : }
244 :
245 0 : return HcclResult::HCCL_SUCCESS;
246 : }
247 :
248 0 : HcclResult InsAlgTemplateBase::PrepBitMask(const u32 queNumPerNeighbor)
249 : {
250 0 : for (auto rankId : tempVTopo_[0]) {
251 : u32 algRank;
252 0 : CHK_RET(GetAlgRank(rankId, tempVTopo_[0], algRank));
253 0 : std::vector<u32> bitPosRank(queNumPerNeighbor);
254 0 : for (u32 posIdx = 0; posIdx < queNumPerNeighbor; posIdx++) {
255 0 : bitPosRank[posIdx] = algRank * queNumPerNeighbor + posIdx;
256 : }
257 0 : std::pair<RankId, std::vector<u32>> newPair(rankId, bitPosRank);
258 0 : rank2BitPos_.insert(newPair);
259 0 : }
260 0 : return HcclResult::HCCL_SUCCESS;
261 : }
262 :
263 0 : std::vector<std::tuple<QId, QId, u32>> InsAlgTemplateBase::CreateMasterSlaveQueNotifiesRequest(u32 queueNum, u32 pairNum,
264 : QId masterId) const
265 : {
266 0 : std::vector<std::tuple<QId, QId, u32>> notifyRequests;
267 0 : HCCL_DEBUG("[Create][MasterSlaveQueNotifiesRequest] queueNum[%u], pairNum[%u], masterId[%u]",
268 : queueNum, pairNum, masterId);
269 0 : if (queueNum == 0 || pairNum == 0) {
270 0 : HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] queueNum or pairNum is zero, "
271 : "return empty notifyRequests");
272 0 : return notifyRequests;
273 : };
274 :
275 0 : u32 slaveNum = queueNum - 1;
276 0 : HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] slavNum[%u]", slaveNum);
277 0 : if (slaveNum < 1 || pairNum < 1) {
278 0 : return notifyRequests;
279 : }
280 0 : notifyRequests.reserve(slaveNum * pairNum);
281 0 : for (QId q = 0; q < queueNum; q++) {
282 0 : if (q == masterId) {
283 0 : continue;
284 : }
285 0 : for (u32 i = 0; i < pairNum; i++) {
286 0 : notifyRequests.emplace_back(std::make_tuple(masterId, q, i));
287 0 : notifyRequests.emplace_back(std::make_tuple(q, masterId, i));
288 : }
289 : }
290 0 : return notifyRequests;
291 0 : }
292 :
293 0 : std::vector<std::tuple<QId, QId, u32>> InsAlgTemplateBase::CreateNotifiesRequestByMap(
294 : std::map<std::tuple<QId, QId>, u32> ¬ifyRequestMap) const
295 : {
296 0 : std::vector<std::tuple<QId, QId, u32>> notifuRequests;
297 :
298 0 : for (auto iter = notifyRequestMap.begin(); iter != notifyRequestMap.end(); iter++) {
299 0 : u32 notifyNum = iter->second;
300 0 : for (u32 i = 0; i < notifyNum; i++) {
301 0 : notifuRequests.emplace_back(std::make_tuple(std::get<0>(iter->first), std::get<1>(iter->first), i));
302 : }
303 : }
304 0 : return notifuRequests;
305 0 : }
306 :
307 0 : std::vector<std::tuple<QId, QId, u32>> InsAlgTemplateBase::MergeNotifiesRequest(
308 : const std::vector<std::vector<std::tuple<QId, QId, u32>>> ¬ifiesRequests) const
309 : {
310 0 : std::vector<std::tuple<QId, QId, u32>> ret;
311 0 : std::map<std::tuple<QId, QId>, u32> requestMap;
312 0 : for (auto ¬ifiesRequest : notifiesRequests) {
313 0 : for (auto &request : notifiesRequest) {
314 0 : QId fromQ = std::get<0>(request);
315 0 : QId toQ = std::get<1>(request);
316 0 : requestMap[std::make_tuple(fromQ, toQ)]++;
317 : }
318 : }
319 0 : return CreateNotifiesRequestByMap(requestMap);
320 0 : }
321 :
322 0 : void InsAlgTemplateBase::SetLoadInfo(const CollAlgParams ¶ms) const
323 : {
324 : (void)params;
325 0 : return;
326 : }
327 :
328 0 : HcclResult InsAlgTemplateBase::GetMaxTransPortDataSize(u64 &maxTransPortDataSize) const
329 : {
330 0 : maxTransPortDataSize = UB_MAX_DATA_SIZE; // 256M
331 0 : return HCCL_SUCCESS;
332 : }
333 :
334 0 : HcclResult InsAlgTemplateBase::CalNumBlocks(u32& numBlocks, u64 dataSize, u32 numBlocksLimit)
335 : {
336 : (void)numBlocks;
337 : (void)dataSize;
338 : (void)numBlocksLimit;
339 0 : HCCL_WARNING("CalNumBlocks not support ins template.");
340 0 : return HCCL_SUCCESS;
341 : }
342 :
343 0 : bool InsAlgTemplateBase::IsPcieLink(const ResLinks &tempLinks) const
344 : {
345 0 : for (auto it = tempLinks.begin(); it != tempLinks.end(); it++) {
346 0 : const std::vector<LinkData>& linkVector = it->second;
347 :
348 0 : for (auto vecIt = linkVector.begin(); vecIt != linkVector.end(); vecIt++) {
349 0 : if (vecIt->GetType() == PortDeploymentType::P2P
350 0 : && vecIt->GetLinkProtocol() == LinkProtocol::PCIE) {
351 0 : HCCL_INFO("IsPcieLink[true]");
352 0 : return true;
353 : }
354 : }
355 : }
356 0 : HCCL_INFO("IsPcieLink[false]");
357 0 : return false;
358 : }
359 0 : HcclResult InsAlgTemplateBase::setPathNumMap(const std::map<u32, u32> &rank2PathNumMap)
360 : {
361 0 : rank2PathNumMap_ = rank2PathNumMap;
362 0 : return HCCL_SUCCESS;
363 : }
364 :
365 : } // namespace Hccl
|