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 "log.h"
12 :
13 : #include "ins_temp_all_to_all_mesh.h"
14 :
15 : namespace Hccl {
16 0 : InsTempAlltoAllMesh::InsTempAlltoAllMesh(const RankId virtualRank, const u32 tempRankSize,
17 : const std::vector<std::vector<RankId>> &tempVTopo,
18 0 : const std::map<RankId, u32> &tempVirtRankMap)
19 0 : : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 : {
21 0 : if (tempRankSize_ == 0) {
22 0 : THROW<InvalidParamsException>(StringFormat("[InsTempAlltoAllMesh] Invalid tempRankSize[%u].", tempRankSize_));
23 : }
24 0 : }
25 :
26 0 : InsTempAlltoAllMesh::~InsTempAlltoAllMesh()
27 : {
28 0 : }
29 :
30 0 : HcclResult InsTempAlltoAllMesh::CalcRes(AlgTempResReq &tempResReq)
31 : {
32 0 : CHK_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq));
33 :
34 0 : auto& linkReq = tempResReq.links;
35 0 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
36 0 : auto remoteRank = resReqIter->first;
37 0 : if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_.at(remoteRank) == 0) {
38 0 : HCCL_ERROR("No path to remoteRank[%d]", remoteRank);
39 0 : return HcclResult::HCCL_E_INTERNAL;
40 : }
41 0 : resReqIter->second = rank2PathNumMap_.at(remoteRank);
42 0 : maxPathNum = std::max(maxPathNum, rank2PathNumMap_.at(remoteRank)); // 每个rank的远端搬运最多需要pathNum条流
43 : }
44 0 : tempResReq.queNum = 1 + maxPathNum * std::min(ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE, tempRankSize_);
45 0 : tempResReq.streamNum = tempResReq.queNum;
46 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
47 :
48 0 : QId centerQ = 0;
49 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
50 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
51 0 : HCCL_DEBUG("[InsTempAlltoAllMesh] Rank[%d], VtopoSize[%zu], requiredQue Num [%u].", myRank_, tempVTopo_[0].size(), tempResReq.queNum);
52 0 : return HcclResult::HCCL_SUCCESS;
53 : }
54 :
55 0 : void InsTempAlltoAllMesh::SetA2ASendRecvInfo(const A2ASendRecvInfo &sendRecvInfo)
56 : {
57 0 : localSendRecvInfo_ = sendRecvInfo;
58 0 : }
59 :
60 0 : HcclResult InsTempAlltoAllMesh::GetScratchBufferInfo(const u64 &scratchBufferSize, DataType dataType)
61 : {
62 : // 需要变更为CCU/AICPU均只加载scratchSize
63 0 : u32 concurrentSendRecvNum = (tempRankSize_ > ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE) ?
64 0 : ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE : tempRankSize_;
65 : // scratch 不需要分两份,userIn直接一步到对端 scratch buffer
66 0 : u64 maxTmpMemSize = scratchBufferSize;
67 0 : u32 typeSize = DataTypeSizeGet(dataType);
68 0 : CHK_PRT_RET(typeSize == 0,
69 : HCCL_ERROR("[InsTempAlltoAllMesh] Rank [%d], Invalid dataSizePerVolume [%u].", myRank_, typeSize),
70 : HcclResult::HCCL_E_INTERNAL);
71 0 : u64 scratchInputMemSize = maxTmpMemSize / (concurrentSendRecvNum * typeSize) * typeSize;
72 0 : CHK_RET(SetBuffBlockSize(scratchInputMemSize));
73 0 : CHK_RET(SetConcurrentSendRecvNum(concurrentSendRecvNum));
74 :
75 0 : HCCL_INFO("[InsTempAlltoAllMesh][GetScratchBufferInfo] concurrentSendRecvNum[%u], scratchInputMemSize[%llu]",
76 : concurrentSendRecvNum, scratchInputMemSize);
77 0 : return HcclResult::HCCL_SUCCESS;
78 : }
79 :
80 0 : HcclResult InsTempAlltoAllMesh::SetBuffBlockSize(const u64 buffBlockSize)
81 : {
82 0 : CHK_PRT_RET(buffBlockSize == 0, HCCL_ERROR("[InsTempAlltoAllMesh][SetBuffBlockSize] buffBlockSize should not be zero"),
83 : HcclResult::HCCL_E_PARA);
84 0 : buffBlockSize_ = buffBlockSize;
85 0 : return HcclResult::HCCL_SUCCESS;
86 : }
87 :
88 0 : HcclResult InsTempAlltoAllMesh::SetConcurrentSendRecvNum(const u32 concurrentSendRecvNum)
89 : {
90 0 : CHK_PRT_RET(concurrentSendRecvNum == 0, HCCL_ERROR("[InsTempAlltoAllMesh][SetConcurrentSendRecvNum] concurrentSendRecvNum should not be zero"),
91 : HcclResult::HCCL_E_PARA);
92 0 : concurrentSendRecvNum_ = concurrentSendRecvNum;
93 0 : return HcclResult::HCCL_SUCCESS;
94 : }
95 :
96 0 : u32 InsTempAlltoAllMesh::CalcStepNum()
97 : {
98 0 : u32 numSubStep = 0;
99 :
100 0 : for (u32 destRank = 0; destRank < tempRankSize_; destRank++) {
101 0 : if (destRank == static_cast<u32>(myRank_)) {
102 0 : continue;
103 : }
104 0 : u32 currRankSendSubStep = ((localSendRecvInfo_.sendLength[destRank] + buffBlockSize_- 1) / buffBlockSize_);
105 0 : u32 currRankRecvSubStep = ((localSendRecvInfo_.recvLength[destRank] + buffBlockSize_- 1) / buffBlockSize_);
106 0 : numSubStep = std::max(numSubStep, std::max(currRankSendSubStep, currRankRecvSubStep));
107 : }
108 0 : return numSubStep;
109 : }
110 :
111 0 : HcclResult InsTempAlltoAllMesh::CalcSendSliceInfo(u32 remoteRank, UsrData &sendSliceInfo)
112 : {
113 : // 判断在打平的视角下,当前的remoteRank 是在本rank的左边还是右边
114 0 : u32 pairNum = concurrentSendRecvNum_ / 2; // 每轮次,会与concurrentSendRecvNum个对端通信,左右各pairNum个
115 0 : u32 gapRight = (tempRankSize_ + remoteRank - myRank_) % tempRankSize_;
116 0 : u32 gapLeft = (tempRankSize_ + myRank_ - remoteRank) % tempRankSize_;
117 0 : u32 sendSrcBuffIdx = 0;
118 0 : u32 sendDstBuffIdx = 0;
119 0 : if (gapLeft < gapRight) {
120 : // 离左边更近
121 0 : u32 gap = gapLeft;
122 0 : sendSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
123 0 : sendDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
124 0 : } else if (gapLeft > gapRight) {
125 : // 离右边更近
126 0 : u32 gap = gapRight;
127 0 : sendSrcBuffIdx = pairNum -1 - ((gap - 1) % pairNum);
128 0 : sendDstBuffIdx = pairNum + ((gap - 1) % pairNum);
129 : } else {
130 0 : sendSrcBuffIdx = 0;
131 0 : sendDstBuffIdx = 0;
132 : }
133 0 : u64 sendDataOffset = 0;
134 0 : u64 remainSendLen = localSendRecvInfo_.sendLength[remoteRank];
135 0 : u64 sendSrcBuffOffset = sendSrcBuffIdx * buffBlockSize_;
136 0 : u64 sendDstBuffOffset = sendDstBuffIdx * buffBlockSize_;
137 0 : while (remainSendLen > 0) {
138 0 : u64 currDataRemainLen = localSendRecvInfo_.sendLength[remoteRank] - sendDataOffset;
139 0 : u64 sendLen = std::min(buffBlockSize_, currDataRemainLen);
140 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[remoteRank] + sendDataOffset;
141 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_] + sendDataOffset;
142 :
143 0 : DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendLen);
144 0 : DataSlice sendScratchSlice = DataSlice(buffInfo_.inBuffType, sendSrcBuffOffset, sendLen);
145 0 : DataSlice sendDstSlice = DataSlice(buffInfo_.outBuffType, sendDstBuffOffset, sendLen);
146 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, sendLen);
147 0 : sendSliceInfo.usrInSlices.emplace_back(userInSlice);
148 0 : sendSliceInfo.scratchInSlices.emplace_back(sendScratchSlice);
149 0 : sendSliceInfo.scratchOutSlices.emplace_back(sendDstSlice);
150 0 : sendSliceInfo.usrOutSlices.emplace_back(userOutSlice);
151 0 : sendDataOffset += sendLen;
152 0 : remainSendLen -= sendLen;
153 : }
154 0 : return HcclResult::HCCL_SUCCESS;
155 : }
156 :
157 0 : HcclResult InsTempAlltoAllMesh::CalcRecvSliceInfo(u32 remoteRank, UsrData &readSliceInfo)
158 : {
159 0 : u32 pairNum = concurrentSendRecvNum_ / 2;
160 0 : u32 gapRight = (remoteRank - myRank_ + tempRankSize_) % tempRankSize_;
161 0 : u32 gapLeft = (myRank_ - remoteRank + tempRankSize_) % tempRankSize_;
162 0 : u32 recvSrcBuffIdx = 0;
163 0 : u32 recvDstBuffIdx = 0;
164 0 : if (gapLeft < gapRight) {
165 0 : u32 gap = gapLeft;
166 0 : recvSrcBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
167 0 : recvDstBuffIdx = pairNum + ((gap - 1) % pairNum);
168 0 : } else if (gapLeft > gapRight) {
169 0 : u32 gap = gapRight;
170 0 : recvSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
171 0 : recvDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
172 : } else {
173 0 : recvSrcBuffIdx = 0;
174 0 : recvDstBuffIdx = 0;
175 : }
176 0 : u64 recvDataOffset = 0;
177 0 : u64 remainRecvLen = localSendRecvInfo_.recvLength[remoteRank];
178 0 : u64 recvSrcBuffOffset = recvSrcBuffIdx * buffBlockSize_;
179 0 : u64 recvDstBuffOffset = recvDstBuffIdx * buffBlockSize_;
180 0 : while(remainRecvLen > 0) {
181 0 : u64 currDataRemainLen = localSendRecvInfo_.recvLength[remoteRank] - recvDataOffset;
182 0 : u64 recvLen = std::min(buffBlockSize_, currDataRemainLen);
183 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[remoteRank] + recvDataOffset;
184 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_] + recvDataOffset;
185 :
186 0 : DataSlice recvInSlice = DataSlice(BufferType::INPUT, userInOffset, recvLen);
187 0 : DataSlice recvScratchSlice = DataSlice(buffInfo_.inBuffType, recvSrcBuffOffset, recvLen);
188 0 : DataSlice recvDstSlice = DataSlice(buffInfo_.outBuffType, recvDstBuffOffset, recvLen);
189 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvLen);
190 0 : readSliceInfo.usrInSlices.emplace_back(recvInSlice);
191 0 : readSliceInfo.scratchInSlices.emplace_back(recvScratchSlice);
192 0 : readSliceInfo.scratchOutSlices.emplace_back(recvDstSlice);
193 0 : readSliceInfo.usrOutSlices.emplace_back(userOutSlice);
194 0 : recvDataOffset += recvLen;
195 0 : remainRecvLen -= recvLen;
196 : }
197 0 : return HcclResult::HCCL_SUCCESS;
198 : }
199 :
200 0 : HcclResult InsTempAlltoAllMesh::CalcSendRecvAllSliceInfo(std::unordered_map<u32, UsrData> &sendSliceInfoMap,
201 : std::unordered_map<u32, UsrData> &recvSliceInfoMap)
202 : {
203 0 : for (u32 remoteRank = 0; remoteRank < tempRankSize_; remoteRank++) {
204 0 : if (remoteRank == static_cast<u32>(myRank_)) {
205 0 : continue;
206 : }
207 0 : UsrData readSliceInfo;
208 0 : UsrData sendSliceInfo;
209 0 : CalcSendSliceInfo(remoteRank, sendSliceInfo);
210 0 : CalcRecvSliceInfo(remoteRank, readSliceInfo);
211 0 : sendSliceInfoMap[remoteRank] = sendSliceInfo;
212 0 : recvSliceInfoMap[remoteRank] = readSliceInfo;
213 0 : }
214 0 : return HcclResult::HCCL_SUCCESS;
215 : }
216 :
217 0 : HcclResult InsTempAlltoAllMesh::CalcCommRankSetforOneLoop(u32 roundIdx, const u32 groupRankSize,
218 : std::vector<u32> &commRanks) const
219 : {
220 0 : commRanks.clear();
221 0 : u32 pairNumPerRound = concurrentSendRecvNum_ / 2;
222 0 : u32 pairSize = (groupRankSize < concurrentSendRecvNum_) ? (groupRankSize + 1) / 2: pairNumPerRound;
223 0 : for (u32 i = roundIdx * pairNumPerRound + 1; i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
224 0 : u32 leftRemoteRank = (myRank_+ tempRankSize_ - i) % tempRankSize_;
225 0 : u32 rightRemoteRank = (myRank_ + i) % tempRankSize_;
226 0 : if (leftRemoteRank == rightRemoteRank) {
227 0 : commRanks.push_back(leftRemoteRank);
228 0 : break;
229 : } else {
230 0 : commRanks.push_back(leftRemoteRank);
231 0 : commRanks.push_back(rightRemoteRank);
232 : }
233 : }
234 0 : return HcclResult::HCCL_SUCCESS;
235 : }
236 :
237 0 : HcclResult InsTempAlltoAllMesh::CopySendDataToScratch(u32 step, const std::vector<u32> &commRanks,
238 : std::unordered_map<u32, UsrData> &sendSliceInfo,
239 : const ResLinks &tempLinks,
240 : std::vector<InsQuePtr> &queues) const
241 : {
242 0 : HCCL_INFO("[InsTempAlltoAllMesh] CopySendDataToScratch");
243 0 : u32 queueId = 0;
244 0 : for(u32 i = 0 ; i < commRanks.size(); i++) {
245 0 : u32 remoteRank = commRanks[i];
246 0 : HCCL_INFO("remoteRank = %u", remoteRank);
247 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
248 0 : std::vector<LinkData>links = tempLinks.at(remoteRank);
249 0 : std::vector<float> dataSplitRate(linkNum);
250 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
251 0 : UsrData &currSendSliceInfo = sendSliceInfo[remoteRank];
252 0 : for(u32 j = 0; j < linkNum; j++)
253 : {
254 0 : InsQuePtr queue = queues[queueId];
255 0 : queueId++;
256 0 : HCCL_INFO("queueId=%u",queueId);
257 0 : if (step < currSendSliceInfo.usrInSlices.size()) {
258 0 : DataSlice &userInSliceAllLinks = currSendSliceInfo.usrInSlices[step];
259 0 : DataSlice &sendScratchSliceAllLinks = currSendSliceInfo.scratchInSlices[step];
260 0 : DataSlice userInSlice = CalcDataSliceForLinks(userInSliceAllLinks, dataSplitRate, j, dataType_);
261 0 : DataSlice sendScratchSlice = CalcDataSliceForLinks(sendScratchSliceAllLinks, dataSplitRate, j, dataType_);
262 :
263 0 : InsQuePtr queue = queues[i];
264 0 : CHK_RET(LocalCopy(queue, userInSlice, sendScratchSlice));
265 0 : }
266 0 : }
267 0 : }
268 0 : return HcclResult::HCCL_SUCCESS;
269 : }
270 :
271 0 : HcclResult InsTempAlltoAllMesh::SendRecvData(u32 step, const std::vector<u32> &commRanks,
272 : std::unordered_map<u32, UsrData> &sendSliceInfo,
273 : std::unordered_map<u32, UsrData> &readSliceInfo,
274 : const ResLinks &tempLinks,
275 : std::vector<InsQuePtr> &queues) const
276 : {
277 0 : CHK_PRT_RET(queues.empty(),
278 : HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] empty queues"), HcclResult::HCCL_E_INTERNAL);
279 0 : CHK_PTR_NULL(queues[0]);
280 :
281 0 : if (commRanks.size() * maxPathNum > queues.size()) {
282 0 : HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] commRanks.size() * maxPathNum > queues.size() is wrong");
283 0 : return HcclResult::HCCL_E_INTERNAL;
284 : }
285 0 : uint64_t queuesId=0;
286 0 : for(u32 i = 0 ; i < commRanks.size(); i++) {
287 0 : s32 remoteRank = static_cast<s32>(commRanks[i]);
288 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
289 0 : UsrData &currSendSliceInfo = sendSliceInfo[remoteRank];
290 0 : UsrData &currReadSliceInfo = readSliceInfo[remoteRank];
291 0 : std::vector<LinkData>links = tempLinks.at(remoteRank);
292 0 : std::vector<float> dataSplitRate(linkNum);
293 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
294 0 : std::vector<DataSlice> &currSendSrcSlices = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.scratchInSlices : currSendSliceInfo.usrInSlices;
295 0 : std::vector<DataSlice> &currSendDstSlices = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.usrOutSlices : currSendSliceInfo.scratchOutSlices;
296 0 : std::vector<DataSlice> &currReadSrcSlices = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.scratchInSlices: currReadSliceInfo.usrInSlices;
297 0 : std::vector<DataSlice> &currReadDstSlices = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.usrOutSlices : currReadSliceInfo.scratchOutSlices;
298 0 : for(u32 j = 0; j < linkNum; j++)
299 : {
300 0 : InsQuePtr queue = queues[queuesId];
301 0 : queuesId++;
302 :
303 0 : LinkData link = tempLinks.at(remoteRank)[j];
304 0 : if (step < currSendSrcSlices.size() && step < currReadSrcSlices.size()) {
305 0 : DataSlice &sendSrcSliceAllLinks = currSendSrcSlices[step];
306 0 : DataSlice &sendDstSliceAllLinks = currSendDstSlices[step];
307 0 : DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
308 0 : DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
309 0 : std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
310 0 : std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
311 0 : SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
312 0 : DataSlice &recvSrcSliceAllLinks = currReadSrcSlices[step];
313 0 : DataSlice &recvDstSliceAllLinks = currReadDstSlices[step];
314 0 : DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
315 0 : DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
316 0 : std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
317 0 : std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
318 0 : SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
319 0 : TxRxSlicesList sendRecvSlice(sendDataSlice, recvDataSlice);
320 :
321 0 : TxRxLinks sendRecvLinks(link, link);
322 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlice);
323 0 : CHK_RET(SendRecv(sendRecvInfo, queue, 0, true, dmaMode_));
324 0 : HCCL_DEBUG("[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send and recv.", step, i, remoteRank);
325 0 : } else if (step < currSendSrcSlices.size()) {
326 0 : DataSlice &sendSrcSliceAllLinks = currSendSrcSlices[step];
327 0 : DataSlice &sendDstSliceAllLinks = currSendDstSlices[step];
328 0 : DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
329 0 : DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
330 0 : std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
331 0 : std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
332 0 : SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
333 0 : DataInfo sendDataInfo(link, sendDataSlice);
334 0 : CHK_RET(Send(sendDataInfo, queue));
335 0 : HCCL_DEBUG("[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send.", step, i, remoteRank);
336 0 : } else if (step < currReadSrcSlices.size()) {
337 0 : DataSlice &recvSrcSliceAllLinks = currReadSrcSlices[step];
338 0 : DataSlice &recvDstSliceAllLinks = currReadDstSlices[step];
339 0 : DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
340 0 : DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
341 0 : std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
342 0 : std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
343 0 : SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
344 0 : DataInfo recvDataInfo(link, recvDataSlice);
345 0 : CHK_RET(Recv(recvDataInfo, queue));
346 0 : HCCL_DEBUG("[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run recv.", step, i, remoteRank);
347 0 : }
348 0 : }
349 0 : }
350 0 : return HcclResult::HCCL_SUCCESS;
351 : }
352 :
353 0 : HcclResult InsTempAlltoAllMesh::CopyRecvDataFromScratch(u32 step, const std::vector<u32> &commRanks,
354 : std::unordered_map<u32, UsrData> &readSliceInfo,
355 : const ResLinks &tempLinks,
356 : std::vector<InsQuePtr> &queues) const
357 : {
358 0 : HCCL_INFO("[InsTempAlltoAllMesh] CopyRecvDataFromScratch");
359 0 : u32 queueId = 0;
360 0 : for(u32 i = 0 ; i < commRanks.size(); i++) {
361 0 : u32 remoteRank = commRanks[i];
362 0 : HCCL_INFO("remoteRank = %u", remoteRank);
363 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
364 0 : std::vector<LinkData>links = tempLinks.at(remoteRank);
365 0 : std::vector<float> dataSplitRate(linkNum);
366 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
367 0 : UsrData &currReadSliceInfo = readSliceInfo[remoteRank];
368 0 : for(u32 j = 0; j < linkNum; j++)
369 : {
370 0 : InsQuePtr queue = queues[queueId];
371 0 : queueId++;
372 0 : HCCL_INFO("queueId=%u",queueId);
373 0 : if (step < currReadSliceInfo.scratchOutSlices.size()) {
374 0 : DataSlice &recvScratchSliceAllLinks = currReadSliceInfo.scratchOutSlices[step];
375 0 : DataSlice &userOutSliceAllLinks = currReadSliceInfo.usrOutSlices[step];
376 0 : DataSlice recvScratchSlice = CalcDataSliceForLinks(recvScratchSliceAllLinks, dataSplitRate, j, dataType_);
377 0 : DataSlice userOutSlice = CalcDataSliceForLinks(userOutSliceAllLinks, dataSplitRate, j, dataType_);
378 0 : CHK_RET(LocalCopy(queue, recvScratchSlice, userOutSlice));
379 : }
380 0 : }
381 0 : }
382 0 : return HcclResult::HCCL_SUCCESS;
383 : }
384 :
385 0 : HcclResult InsTempAlltoAllMesh::RunSendRecvBufferLoop(u32 step, const std::vector<u32> &commRanks,
386 : std::unordered_map<u32, UsrData> &sendSliceInfoMap,
387 : std::unordered_map<u32, UsrData> &recvSliceInfoMap,
388 : const ResLinks &tempLinks,
389 : std::vector<InsQuePtr> &queues) const
390 : {
391 0 : if (commRanks.size() > 1) {
392 0 : CHK_RET(PreSyncInterQueues(queues));
393 : }
394 :
395 0 : if (dmaMode_ == DmaMode::PUT) {
396 0 : CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
397 0 : CHK_RET(CopyRecvDataFromScratch(step, commRanks, recvSliceInfoMap, tempLinks, queues));
398 0 : } else if (dmaMode_ == DmaMode::GET) {
399 0 : CHK_RET(CopySendDataToScratch(step, commRanks, sendSliceInfoMap, tempLinks, queues));
400 0 : CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
401 : }
402 :
403 0 : if (commRanks.size() > 1) {
404 0 : CHK_RET(PostSyncInterQueues(queues));
405 : }
406 0 : return HcclResult::HCCL_SUCCESS;
407 : }
408 :
409 0 : HcclResult InsTempAlltoAllMesh::RunSendRecvForAllRanks(u32 step, std::unordered_map<u32, UsrData> &sendSliceInfoMap,
410 : std::unordered_map<u32, UsrData> &recvSliceInfoMap, const ResLinks &tempLinks, std::vector<InsQuePtr> &queues)
411 : {
412 0 : u64 commLoops = (tempRankSize_ + concurrentSendRecvNum_ - 1) / concurrentSendRecvNum_;
413 0 : u32 leftRankSize = tempRankSize_ - 1; // leftRankSize中去掉本卡
414 0 : std::vector<u32> commRanks;
415 0 : for (u32 roundIdx = 0; roundIdx < commLoops && leftRankSize > 0; roundIdx++) {
416 0 : u32 groupRankSize = (leftRankSize > concurrentSendRecvNum_) ? concurrentSendRecvNum_ : leftRankSize;
417 0 : CHK_RET(CalcCommRankSetforOneLoop(roundIdx, groupRankSize, commRanks));
418 0 : CHK_RET(RunSendRecvBufferLoop(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
419 0 : leftRankSize -= groupRankSize;
420 0 : HCCL_DEBUG("[InsTempAlltoAllMesh][RunSendRecvForAllRanks] commRanksSize[%zu], roundIdx[%u] finish.", commRanks.size(), roundIdx);
421 : }
422 0 : return HcclResult::HCCL_SUCCESS;
423 0 : }
424 :
425 0 : HcclResult InsTempAlltoAllMesh::LocalDataCopy(InsQuePtr tempInsQue)
426 : {
427 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_];
428 0 : u64 sendSize = localSendRecvInfo_.sendLength[myRank_];
429 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_];
430 0 : u64 recvSize = localSendRecvInfo_.recvLength[myRank_];
431 0 : DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendSize);
432 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvSize);
433 0 : CHK_RET(LocalCopy(tempInsQue, userInSlice, userOutSlice));
434 0 : return HcclResult::HCCL_SUCCESS;
435 : }
436 :
437 0 : HcclResult InsTempAlltoAllMesh::Run(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
438 : const BuffInfo &buffInfo,
439 : const ResLinks &tempLinks,
440 : std::vector<InsQuePtr> &tempInsQues)
441 : {
442 : (void) tempFuncs;
443 : (void) sliceInfoVec;
444 0 : if (tempInsQues.size() == 0) {
445 0 : HCCL_ERROR("[CcuTempAlltoAllMesh1D] tempInsQues.size() is zero.");
446 0 : return HcclResult::HCCL_E_PARA;
447 : }
448 0 : dmaMode_ = DmaMode::PUT;
449 0 : if (IsPcieLink(tempLinks)) {
450 0 : dmaMode_ = DmaMode::GET;
451 : }
452 0 : buffInfo_ = buffInfo;
453 0 : u32 totalStep = CalcStepNum();
454 0 : HCCL_INFO("[InsTempAlltoAllMesh] AllToAll full mesh total step is [%u], tempInsQuesSize[%zu].", totalStep, tempInsQues.size());
455 0 : std::unordered_map<u32, UsrData> sendSliceInfoMap;
456 0 : std::unordered_map<u32, UsrData> recvSliceInfoMap;
457 0 : CHK_RET(CalcSendRecvAllSliceInfo(sendSliceInfoMap, recvSliceInfoMap));
458 0 : std::vector<InsQuePtr> localCopyQues;
459 0 : if (tempInsQues.size() > 1) {
460 0 : localCopyQues.push_back(tempInsQues[0]);
461 0 : localCopyQues.push_back(tempInsQues.back());
462 : }
463 :
464 0 : if (localCopyQues.size() > 1) {
465 0 : CHK_RET(PreSyncInterQueues(localCopyQues));
466 : }
467 : // 最后一条流做localCopy
468 0 : CHK_RET(LocalDataCopy(tempInsQues.back()));
469 : // 其余的流去做sendRecv
470 0 : std::vector<InsQuePtr> sendRecvQues(tempInsQues.begin(), tempInsQues.begin() + (tempInsQues.size() - 1));
471 0 : for (u32 step = 0; step < totalStep; step++) {
472 0 : CHK_RET(RunSendRecvForAllRanks(step, sendSliceInfoMap, recvSliceInfoMap, tempLinks, sendRecvQues));
473 0 : HCCL_DEBUG("[InsTempAlltoAllMesh] AllToAll full mesh step[%u] execute success", step);
474 : }
475 0 : if (localCopyQues.size() > 1) {
476 0 : CHK_RET(PostSyncInterQueues(localCopyQues));
477 : }
478 0 : HCCL_INFO("[InsTempAlltoAllMesh] AllToAll full mesh rank[%d] finish.", myRank_);
479 0 : return HcclResult::HCCL_SUCCESS;
480 0 : }
481 :
482 : } // namespace Hccl
|