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 0 : InsAlgTemplateBase::InsAlgTemplateBase(
17 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
18 0 : const std::map<RankId, u32>& tempVirtRankMap)
19 0 : : myRank_(virtualRank),
20 0 : tempRankSize_(tempRankSize),
21 0 : tempVTopo_(tempVTopo),
22 0 : tempVirtRankMap_(tempVirtRankMap)
23 0 : {}
24 :
25 0 : InsAlgTemplateBase::~InsAlgTemplateBase() {}
26 :
27 0 : void InsAlgTemplateBase::SetCollOp(const CollAlgOperator& op)
28 : {
29 0 : op_ = op;
30 0 : return;
31 : }
32 :
33 0 : void InsAlgTemplateBase::SetDmaMode(const DmaMode dmaMode)
34 : {
35 0 : dmaMode_ = dmaMode;
36 0 : return;
37 : }
38 :
39 0 : void InsAlgTemplateBase::SetRoot(const u32 root)
40 : {
41 0 : root_ = root;
42 0 : return;
43 : }
44 :
45 0 : u64 InsAlgTemplateBase::CalcLoopMaxCount(ParamPool& paramPool)
46 : {
47 0 : u64 loopMaxCount = 0;
48 0 : if (paramPool.params.opMode == OpMode::OPBASE) {
49 : u64 maxLoopSize
50 0 : = 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 : HcclResult
81 0 : InsAlgTemplateBase::CalcSliceInfo(const AllignInfo& allignInfo, const u64 dataSize, 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(
126 : const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, 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(
228 : PreSync(queIdx, syncInsQues) != HcclResult::HCCL_SUCCESS,
229 : HCCL_ERROR(
230 : "[InsCollAlgFactory] Rank [%d], Que [%u], Semaphore Synchronization Failed.", myRank_,
231 : syncInsQues[queIdx]->GetId()),
232 : HcclResult::HCCL_E_INTERNAL);
233 : }
234 :
235 0 : return HcclResult::HCCL_SUCCESS;
236 : }
237 :
238 0 : HcclResult InsAlgTemplateBase::PostSyncInterQueues(std::vector<InsQuePtr>& syncInsQues) const
239 : {
240 0 : for (size_t queIdx = 0; queIdx < syncInsQues.size(); queIdx++) {
241 0 : CHK_PRT_RET(
242 : PostSync(queIdx, syncInsQues) != HcclResult::HCCL_SUCCESS,
243 : HCCL_ERROR(
244 : "[InsCollAlgFactory] Rank [%d], Que [%u], Semaphore Synchronization Failed.", myRank_,
245 : syncInsQues[queIdx]->GetId()),
246 : HcclResult::HCCL_E_INTERNAL);
247 : }
248 :
249 0 : return HcclResult::HCCL_SUCCESS;
250 : }
251 :
252 0 : HcclResult InsAlgTemplateBase::PrepBitMask(const u32 queNumPerNeighbor)
253 : {
254 0 : for (auto rankId : tempVTopo_[0]) {
255 : u32 algRank;
256 0 : CHK_RET(GetAlgRank(rankId, tempVTopo_[0], algRank));
257 0 : std::vector<u32> bitPosRank(queNumPerNeighbor);
258 0 : for (u32 posIdx = 0; posIdx < queNumPerNeighbor; posIdx++) {
259 0 : bitPosRank[posIdx] = algRank * queNumPerNeighbor + posIdx;
260 : }
261 0 : std::pair<RankId, std::vector<u32>> newPair(rankId, bitPosRank);
262 0 : rank2BitPos_.insert(newPair);
263 0 : }
264 0 : return HcclResult::HCCL_SUCCESS;
265 : }
266 :
267 : std::vector<std::tuple<QId, QId, u32>>
268 0 : InsAlgTemplateBase::CreateMasterSlaveQueNotifiesRequest(u32 queueNum, u32 pairNum, QId masterId) const
269 : {
270 0 : std::vector<std::tuple<QId, QId, u32>> notifyRequests;
271 0 : HCCL_DEBUG(
272 : "[Create][MasterSlaveQueNotifiesRequest] queueNum[%u], pairNum[%u], masterId[%u]", queueNum, pairNum, masterId);
273 0 : if (queueNum == 0 || pairNum == 0) {
274 0 : HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] queueNum or pairNum is zero, "
275 : "return empty notifyRequests");
276 0 : return notifyRequests;
277 : };
278 :
279 0 : u32 slaveNum = queueNum - 1;
280 0 : HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] slavNum[%u]", slaveNum);
281 0 : if (slaveNum < 1 || pairNum < 1) {
282 0 : return notifyRequests;
283 : }
284 0 : notifyRequests.reserve(slaveNum * pairNum);
285 0 : for (QId q = 0; q < queueNum; q++) {
286 0 : if (q == masterId) {
287 0 : continue;
288 : }
289 0 : for (u32 i = 0; i < pairNum; i++) {
290 0 : notifyRequests.emplace_back(std::make_tuple(masterId, q, i));
291 0 : notifyRequests.emplace_back(std::make_tuple(q, masterId, i));
292 : }
293 : }
294 0 : return notifyRequests;
295 0 : }
296 :
297 : std::vector<std::tuple<QId, QId, u32>>
298 0 : InsAlgTemplateBase::CreateNotifiesRequestByMap(std::map<std::tuple<QId, QId>, u32>& notifyRequestMap) const
299 : {
300 0 : std::vector<std::tuple<QId, QId, u32>> notifuRequests;
301 :
302 0 : for (auto iter = notifyRequestMap.begin(); iter != notifyRequestMap.end(); iter++) {
303 0 : u32 notifyNum = iter->second;
304 0 : for (u32 i = 0; i < notifyNum; i++) {
305 0 : notifuRequests.emplace_back(std::make_tuple(std::get<0>(iter->first), std::get<1>(iter->first), i));
306 : }
307 : }
308 0 : return notifuRequests;
309 0 : }
310 :
311 0 : std::vector<std::tuple<QId, QId, u32>> InsAlgTemplateBase::MergeNotifiesRequest(
312 : const std::vector<std::vector<std::tuple<QId, QId, u32>>>& notifiesRequests) const
313 : {
314 0 : std::vector<std::tuple<QId, QId, u32>> ret;
315 0 : std::map<std::tuple<QId, QId>, u32> requestMap;
316 0 : for (auto& notifiesRequest : notifiesRequests) {
317 0 : for (auto& request : notifiesRequest) {
318 0 : QId fromQ = std::get<0>(request);
319 0 : QId toQ = std::get<1>(request);
320 0 : requestMap[std::make_tuple(fromQ, toQ)]++;
321 : }
322 : }
323 0 : return CreateNotifiesRequestByMap(requestMap);
324 0 : }
325 :
326 0 : void InsAlgTemplateBase::SetLoadInfo(const CollAlgParams& params) const
327 : {
328 : (void)params;
329 0 : return;
330 : }
331 :
332 0 : HcclResult InsAlgTemplateBase::GetMaxTransPortDataSize(u64& maxTransPortDataSize) const
333 : {
334 0 : maxTransPortDataSize = UB_MAX_DATA_SIZE; // 256M
335 0 : return HCCL_SUCCESS;
336 : }
337 :
338 0 : HcclResult InsAlgTemplateBase::CalNumBlocks(u32& numBlocks, u64 dataSize, u32 numBlocksLimit)
339 : {
340 : (void)numBlocks;
341 : (void)dataSize;
342 : (void)numBlocksLimit;
343 0 : HCCL_WARNING("CalNumBlocks not support ins template.");
344 0 : return HCCL_SUCCESS;
345 : }
346 :
347 0 : bool InsAlgTemplateBase::IsPcieLink(const ResLinks& tempLinks) const
348 : {
349 0 : for (auto it = tempLinks.begin(); it != tempLinks.end(); it++) {
350 0 : const std::vector<LinkData>& linkVector = it->second;
351 :
352 0 : for (auto vecIt = linkVector.begin(); vecIt != linkVector.end(); vecIt++) {
353 0 : if (vecIt->GetType() == PortDeploymentType::P2P && vecIt->GetLinkProtocol() == LinkProtocol::PCIE) {
354 0 : HCCL_INFO("IsPcieLink[true]");
355 0 : return true;
356 : }
357 : }
358 : }
359 0 : HCCL_INFO("IsPcieLink[false]");
360 0 : return false;
361 : }
362 0 : HcclResult InsAlgTemplateBase::setPathNumMap(const std::map<u32, u32>& rank2PathNumMap)
363 : {
364 0 : rank2PathNumMap_ = rank2PathNumMap;
365 0 : return HCCL_SUCCESS;
366 : }
367 :
368 : } // namespace Hccl
|